mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
updates
This commit is contained in:
@@ -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
@@ -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
|
import argparse
|
||||||
|
|
||||||
|
|
||||||
@@ -33,6 +39,11 @@ def _get_model_args(parser: argparse.ArgumentParser) -> None:
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="If we should cast DiT params to a lower precision.",
|
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:
|
def _get_dataset_args(parser: argparse.ArgumentParser) -> None:
|
||||||
@@ -93,6 +104,18 @@ def _get_validation_args(parser: argparse.ArgumentParser) -> None:
|
|||||||
default=50,
|
default=50,
|
||||||
help="Run validation every X training steps. Validation consists of running the validation prompt `args.num_validation_videos` times.",
|
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(
|
parser.add_argument(
|
||||||
"--enable_model_cpu_offload",
|
"--enable_model_cpu_offload",
|
||||||
action="store_true",
|
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"],
|
default=["to_k", "to_q", "to_v", "to_out.0"],
|
||||||
help="Target modules to train LoRA for.",
|
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(
|
parser.add_argument(
|
||||||
"--output_dir",
|
"--output_dir",
|
||||||
type=str,
|
type=str,
|
||||||
@@ -163,37 +175,6 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
|
|||||||
default=None,
|
default=None,
|
||||||
help="Total number of training steps to perform. If provided, overrides `--num_train_epochs`.",
|
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(
|
parser.add_argument(
|
||||||
"--gradient_checkpointing",
|
"--gradient_checkpointing",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
@@ -210,45 +191,12 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
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(
|
parser.add_argument(
|
||||||
"--lr_warmup_steps",
|
"--lr_warmup_steps",
|
||||||
type=int,
|
type=int,
|
||||||
default=200,
|
default=200,
|
||||||
help="Number of steps for the warmup in the lr scheduler.",
|
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:
|
def _get_optimizer_args(parser: argparse.ArgumentParser) -> None:
|
||||||
@@ -256,78 +204,15 @@ def _get_optimizer_args(parser: argparse.ArgumentParser) -> None:
|
|||||||
"--optimizer",
|
"--optimizer",
|
||||||
type=lambda s: s.lower(),
|
type=lambda s: s.lower(),
|
||||||
default="adam",
|
default="adam",
|
||||||
choices=["adam", "adamw", "prodigy", "came"],
|
choices=["adam", "adamw"],
|
||||||
help=("The optimizer type to use."),
|
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(
|
parser.add_argument(
|
||||||
"--weight_decay",
|
"--weight_decay",
|
||||||
type=float,
|
type=float,
|
||||||
default=0.01,
|
default=0.01,
|
||||||
help="Weight decay to use for optimizer.",
|
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:
|
def _get_configuration_args(parser: argparse.ArgumentParser) -> None:
|
||||||
@@ -349,12 +234,6 @@ def _get_configuration_args(parser: argparse.ArgumentParser) -> None:
|
|||||||
default=None,
|
default=None,
|
||||||
help="The name of the repository to keep in sync with the local `output_dir`.",
|
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(
|
parser.add_argument(
|
||||||
"--allow_tf32",
|
"--allow_tf32",
|
||||||
action="store_true",
|
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"
|
" 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(
|
parser.add_argument(
|
||||||
"--report_to",
|
"--report_to",
|
||||||
type=str,
|
type=str,
|
||||||
default=None,
|
default=None,
|
||||||
help=(
|
help="If logging to wandb."
|
||||||
'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.'
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
|
||||||
@@ -1,9 +1,11 @@
|
|||||||
#!/bin/bash
|
#!/bin/bash
|
||||||
|
|
||||||
GPU_ID=0
|
GPU_ID=0
|
||||||
VIDEO_DIR=/home/sayak/cogvideox-factory/video-dataset-disney-organized
|
VIDEO_DIR=video-dataset-disney-organized
|
||||||
OUTPUT_DIR=videos_prepared
|
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
|
CUDA_VISIBLE_DEVICES=$GPU_ID python embed.py $OUTPUT_DIR --shape=37x480x848
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
peft
|
||||||
|
transformers
|
||||||
|
wandb
|
||||||
|
torch
|
||||||
|
torchvision
|
||||||
|
moviepy
|
||||||
|
click
|
||||||
@@ -16,41 +16,22 @@
|
|||||||
import gc
|
import gc
|
||||||
import random
|
import random
|
||||||
from glob import glob
|
from glob import glob
|
||||||
import logging
|
|
||||||
import math
|
import math
|
||||||
import os
|
import os
|
||||||
import shutil
|
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from datetime import timedelta
|
import numpy as np
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Dict, Tuple, List
|
from typing import Any, Dict, Tuple, List
|
||||||
|
|
||||||
import diffusers
|
|
||||||
import torch
|
import torch
|
||||||
import transformers
|
|
||||||
import wandb
|
import wandb
|
||||||
from accelerate import Accelerator, DistributedType
|
from diffusers import FlowMatchEulerDiscreteScheduler, MochiPipeline, MochiTransformer3DModel
|
||||||
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.models.autoencoders.vae import DiagonalGaussianDistribution
|
||||||
from diffusers.optimization import get_scheduler
|
|
||||||
from diffusers.training_utils import cast_training_params
|
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.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 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 torch.utils.data import DataLoader
|
||||||
from tqdm.auto import tqdm
|
from tqdm.auto import tqdm
|
||||||
|
|
||||||
@@ -63,10 +44,23 @@ import sys
|
|||||||
|
|
||||||
sys.path.append("..")
|
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(
|
def save_model_card(
|
||||||
@@ -84,7 +78,7 @@ def save_model_card(
|
|||||||
widget_dict.append(
|
widget_dict.append(
|
||||||
{
|
{
|
||||||
"text": validation_prompt if validation_prompt else " ",
|
"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(
|
def log_validation(
|
||||||
accelerator: Accelerator,
|
|
||||||
pipe: MochiPipeline,
|
pipe: MochiPipeline,
|
||||||
args: Dict[str, Any],
|
args: Dict[str, Any],
|
||||||
pipeline_args: Dict[str, Any],
|
pipeline_args: Dict[str, Any],
|
||||||
epoch,
|
epoch,
|
||||||
|
wandb_run: str = None,
|
||||||
is_final_validation: bool = False,
|
is_final_validation: bool = False,
|
||||||
):
|
):
|
||||||
logger.info(
|
print(
|
||||||
f"Running validation... \n Generating {args.num_validation_videos} videos with prompt: {pipeline_args['prompt']}."
|
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:
|
if not args.enable_model_cpu_offload:
|
||||||
pipe = pipe.to(accelerator.device)
|
pipe = pipe.to("cuda")
|
||||||
|
|
||||||
# run inference
|
# 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 = []
|
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):
|
for _ in range(args.num_validation_videos):
|
||||||
video = pipe(**pipeline_args, generator=generator, output_type="np").frames[0]
|
video = pipe(**pipeline_args, generator=generator, output_type="np").frames[0]
|
||||||
videos.append(video)
|
videos.append(video)
|
||||||
|
|
||||||
for tracker in accelerator.trackers:
|
video_filenames = []
|
||||||
phase_name = "test" if is_final_validation else "validation"
|
for i, video in enumerate(videos):
|
||||||
if tracker.name == "wandb":
|
prompt = (
|
||||||
video_filenames = []
|
pipeline_args["prompt"][:25]
|
||||||
for i, video in enumerate(videos):
|
.replace(" ", "_")
|
||||||
prompt = (
|
.replace(" ", "_")
|
||||||
pipeline_args["prompt"][:25]
|
.replace("'", "_")
|
||||||
.replace(" ", "_")
|
.replace('"', "_")
|
||||||
.replace(" ", "_")
|
.replace("/", "_")
|
||||||
.replace("'", "_")
|
)
|
||||||
.replace('"', "_")
|
filename = os.path.join(args.output_dir, f"{phase_name}_video_{i}_{prompt}.mp4")
|
||||||
.replace("/", "_")
|
export_to_video(video, filename, fps=30)
|
||||||
)
|
video_filenames.append(filename)
|
||||||
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(
|
if wandb_run:
|
||||||
{
|
wandb.log(
|
||||||
phase_name: [
|
{
|
||||||
wandb.Video(filename, caption=f"{i}: {pipeline_args['prompt']}", fps=30)
|
phase_name: [
|
||||||
for i, filename in enumerate(video_filenames)
|
wandb.Video(filename, caption=f"{i}: {pipeline_args['prompt']}", fps=30)
|
||||||
]
|
for i, filename in enumerate(video_filenames)
|
||||||
}
|
]
|
||||||
)
|
}
|
||||||
|
)
|
||||||
|
|
||||||
return videos
|
return videos
|
||||||
|
|
||||||
@@ -231,63 +224,24 @@ class CollateFunction:
|
|||||||
|
|
||||||
|
|
||||||
def main(args):
|
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:
|
if args.report_to == "wandb" and args.hub_token is not None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token."
|
"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."
|
" 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
|
# Handle the repository creation
|
||||||
if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
|
if args.output_dir is not None:
|
||||||
if args.output_dir is not None:
|
os.makedirs(args.output_dir, exist_ok=True)
|
||||||
os.makedirs(args.output_dir, exist_ok=True)
|
|
||||||
|
|
||||||
if args.push_to_hub:
|
if args.push_to_hub:
|
||||||
repo_id = create_repo(
|
repo_id = create_repo(
|
||||||
repo_id=args.hub_model_id or Path(args.output_dir).name,
|
repo_id=args.hub_model_id or Path(args.output_dir).name,
|
||||||
exist_ok=True,
|
exist_ok=True,
|
||||||
).repo_id
|
).repo_id
|
||||||
|
|
||||||
# Prepare models and scheduler
|
# Prepare models and scheduler
|
||||||
transformer = MochiTransformer3DModel.from_pretrained(
|
transformer = MochiTransformer3DModel.from_pretrained(
|
||||||
@@ -300,44 +254,14 @@ def main(args):
|
|||||||
args.pretrained_model_name_or_path, subfolder="scheduler"
|
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.requires_grad_(False)
|
||||||
transformer.to(accelerator.device)
|
transformer.to("cuda")
|
||||||
if args.gradient_checkpointing:
|
if args.gradient_checkpointing:
|
||||||
transformer.enable_gradient_checkpointing()
|
transformer.enable_gradient_checkpointing()
|
||||||
if args.cast_dit:
|
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
|
# now we will add new LoRA weights to the attention layers
|
||||||
transformer_lora_config = LoraConfig(
|
transformer_lora_config = LoraConfig(
|
||||||
@@ -348,131 +272,25 @@ def main(args):
|
|||||||
)
|
)
|
||||||
transformer.add_adapter(transformer_lora_config)
|
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,
|
# Enable TF32 for faster training on Ampere GPUs,
|
||||||
# cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
|
# cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
|
||||||
if args.allow_tf32 and torch.cuda.is_available():
|
if args.allow_tf32 and torch.cuda.is_available():
|
||||||
torch.backends.cuda.matmul.allow_tf32 = True
|
torch.backends.cuda.matmul.allow_tf32 = True
|
||||||
|
|
||||||
if args.scale_lr:
|
if args.scale_lr:
|
||||||
args.learning_rate = (
|
args.learning_rate = args.learning_rate * args.train_batch_size
|
||||||
args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes
|
|
||||||
)
|
|
||||||
# only upcast trainable parameters (LoRA) into fp32
|
# only upcast trainable parameters (LoRA) into fp32
|
||||||
cast_training_params([transformer], dtype=torch.float32)
|
cast_training_params([transformer], dtype=torch.float32)
|
||||||
|
|
||||||
|
# Prepare optimizer
|
||||||
transformer_lora_parameters = list(filter(lambda p: p.requires_grad, transformer.parameters()))
|
transformer_lora_parameters = list(filter(lambda p: p.requires_grad, transformer.parameters()))
|
||||||
|
num_trainable_parameters = sum(param.numel() for param in transformer_lora_parameters)
|
||||||
# Optimization parameters
|
optimizer = torch.optim.AdamW(transformer_lora_parameters, lr=args.learning_rate, weight_decay=args.weight_decay)
|
||||||
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.")
|
|
||||||
|
|
||||||
# Dataset and DataLoader
|
# Dataset and DataLoader
|
||||||
train_vids = list(sorted(glob(f"{args.data_root}/*.mp4")))
|
train_vids = list(sorted(glob(f"{args.data_root}/*.mp4")))
|
||||||
train_vids = [v for v in train_vids if not v.endswith(".recon.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}"
|
assert len(train_vids) > 0, f"No training data found in {args.data_root}"
|
||||||
|
|
||||||
collate_fn = CollateFunction(caption_dropout=args.caption_dropout)
|
collate_fn = CollateFunction(caption_dropout=args.caption_dropout)
|
||||||
@@ -485,46 +303,19 @@ def main(args):
|
|||||||
pin_memory=args.pin_memory,
|
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
|
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:
|
if args.max_train_steps is None:
|
||||||
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
|
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
|
||||||
overrode_max_train_steps = True
|
overrode_max_train_steps = True
|
||||||
|
|
||||||
if args.use_cpu_offload_optimizer:
|
lr_scheduler = get_cosine_annealing_lr_scheduler(
|
||||||
lr_scheduler = None
|
optimizer, warmup_steps=args.lr_warmup_steps, total_steps=args.max_train_steps
|
||||||
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.
|
# 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:
|
if overrode_max_train_steps:
|
||||||
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
|
args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
|
||||||
# Afterwards we recalculate our number of training epochs
|
# 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.
|
# We need to initialize the trackers we use, and also store our configuration.
|
||||||
# The trackers initializes automatically on the main process.
|
# 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"
|
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 =====")
|
print("===== Memory before training =====")
|
||||||
reset_memory(accelerator.device)
|
reset_memory("cuda")
|
||||||
print_memory(accelerator.device)
|
print_memory("cuda")
|
||||||
|
|
||||||
# Train!
|
# 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
|
global_step = 0
|
||||||
first_epoch = 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(
|
progress_bar = tqdm(
|
||||||
range(0, args.max_train_steps),
|
range(0, args.max_train_steps),
|
||||||
initial=initial_global_step,
|
initial=global_step,
|
||||||
desc="Steps",
|
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):
|
for epoch in range(first_epoch, args.num_train_epochs):
|
||||||
transformer.train()
|
transformer.train()
|
||||||
|
|
||||||
for step, batch in enumerate(train_dataloader):
|
for step, batch in enumerate(train_dataloader):
|
||||||
models_to_accumulate = [transformer]
|
with torch.no_grad():
|
||||||
|
z = batch["z"].to("cuda")
|
||||||
with accelerator.accumulate(models_to_accumulate):
|
eps = batch["eps"].to("cuda")
|
||||||
z = batch["z"]
|
sigma = batch["sigma"].to("cuda")
|
||||||
# revisit
|
prompt_embeds = batch["prompt_embeds"].to("cuda")
|
||||||
# if has_latents_mean and has_latents_std:
|
prompt_attention_mask = batch["prompt_attention_mask"].to("cuda")
|
||||||
# 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"]
|
|
||||||
|
|
||||||
sigma_bcthw = sigma[:, None, None, None, None] # [B, 1, 1, 1, 1]
|
sigma_bcthw = sigma[:, None, None, None, None] # [B, 1, 1, 1, 1]
|
||||||
# Add noise according to flow matching.
|
# Add noise according to flow matching.
|
||||||
@@ -613,80 +367,35 @@ def main(args):
|
|||||||
z_sigma = (1 - sigma_bcthw) * z + sigma_bcthw * eps
|
z_sigma = (1 - sigma_bcthw) * z + sigma_bcthw * eps
|
||||||
ut = z - eps
|
ut = z - eps
|
||||||
|
|
||||||
# Predict the noise residual
|
|
||||||
# (1 - sigma) because of
|
# (1 - sigma) because of
|
||||||
# https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/src/genmo/mochi_preview/dit/joint_model/asymm_models_joint.py#L656
|
# 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.
|
# Also, we operate on the scaled version of the `timesteps` directly in the `diffusers` implementation.
|
||||||
timesteps = (1 - sigma) * scheduler.config.num_train_timesteps
|
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:
|
with torch.autocast("cuda", torch.bfloat16):
|
||||||
# no grad norm for now, following the original code
|
model_pred = transformer(
|
||||||
# https://github.com/genmoai/mochi/blob/aba74c1b5e0755b1fa3343d9e4bd22e89de77ab1/demos/fine_tuner/train.py#L380
|
hidden_states=z_sigma,
|
||||||
# gradient_norm_before_clip = get_gradient_norm(transformer_lora_parameters)
|
encoder_hidden_states=prompt_embeds,
|
||||||
# accelerator.clip_grad_norm_(transformer_lora_parameters, args.max_grad_norm)
|
encoder_attention_mask=prompt_attention_mask,
|
||||||
# gradient_norm_after_clip = get_gradient_norm(transformer_lora_parameters)
|
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.step()
|
optimizer.zero_grad()
|
||||||
optimizer.zero_grad()
|
lr_scheduler.step()
|
||||||
|
|
||||||
if not args.use_cpu_offload_optimizer:
|
progress_bar.update(1)
|
||||||
lr_scheduler.step()
|
global_step += 1
|
||||||
|
|
||||||
# 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
|
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}
|
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)
|
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:
|
if global_step >= args.max_train_steps:
|
||||||
break
|
break
|
||||||
@@ -694,82 +403,15 @@ def main(args):
|
|||||||
if global_step >= args.max_train_steps:
|
if global_step >= args.max_train_steps:
|
||||||
break
|
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:
|
||||||
if args.validation_prompt is not None and (epoch + 1) % args.validation_epochs == 0:
|
print("===== Memory before validation =====")
|
||||||
accelerator.print("===== Memory before validation =====")
|
print_memory("cuda")
|
||||||
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)
|
|
||||||
|
|
||||||
|
transformer.eval()
|
||||||
pipe = MochiPipeline.from_pretrained(
|
pipe = MochiPipeline.from_pretrained(
|
||||||
args.pretrained_model_name_or_path,
|
args.pretrained_model_name_or_path,
|
||||||
|
transformer=transformer,
|
||||||
|
scheduler=scheduler,
|
||||||
revision=args.revision,
|
revision=args.revision,
|
||||||
variant=args.variant,
|
variant=args.variant,
|
||||||
)
|
)
|
||||||
@@ -781,12 +423,6 @@ def main(args):
|
|||||||
if args.enable_model_cpu_offload:
|
if args.enable_model_cpu_offload:
|
||||||
pipe.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)
|
validation_prompts = args.validation_prompt.split(args.validation_prompt_separator)
|
||||||
for validation_prompt in validation_prompts:
|
for validation_prompt in validation_prompts:
|
||||||
pipeline_args = {
|
pipeline_args = {
|
||||||
@@ -797,40 +433,104 @@ def main(args):
|
|||||||
"width": args.width,
|
"width": args.width,
|
||||||
"max_sequence_length": 256,
|
"max_sequence_length": 256,
|
||||||
}
|
}
|
||||||
|
log_validation(
|
||||||
video = log_validation(
|
|
||||||
accelerator=accelerator,
|
|
||||||
pipe=pipe,
|
pipe=pipe,
|
||||||
args=args,
|
args=args,
|
||||||
pipeline_args=pipeline_args,
|
pipeline_args=pipeline_args,
|
||||||
epoch=epoch,
|
epoch=epoch,
|
||||||
is_final_validation=True,
|
wandb_run=wandb_run,
|
||||||
)
|
)
|
||||||
validation_outputs.extend(video)
|
|
||||||
|
|
||||||
accelerator.print("===== Memory after testing =====")
|
print("===== Memory after validation =====")
|
||||||
print_memory(accelerator.device)
|
print_memory("cuda")
|
||||||
reset_memory(accelerator.device)
|
reset_memory("cuda")
|
||||||
torch.cuda.synchronize(accelerator.device)
|
|
||||||
|
|
||||||
if args.push_to_hub:
|
del pipe.text_encoder
|
||||||
save_model_card(
|
del pipe.vae
|
||||||
repo_id,
|
del pipe
|
||||||
videos=validation_outputs,
|
gc.collect()
|
||||||
base_model=args.pretrained_model_name_or_path,
|
torch.cuda.empty_cache()
|
||||||
validation_prompt=args.validation_prompt,
|
|
||||||
repo_folder=args.output_dir,
|
transformer.train()
|
||||||
fps=args.fps,
|
|
||||||
|
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(
|
validation_outputs.extend(video)
|
||||||
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}.")
|
|
||||||
|
|
||||||
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__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -2,38 +2,35 @@
|
|||||||
export NCCL_P2P_DISABLE=1
|
export NCCL_P2P_DISABLE=1
|
||||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
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"
|
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 \
|
--pretrained_model_name_or_path $MODEL \
|
||||||
|
--cast_dit \
|
||||||
--data_root $DATA_ROOT \
|
--data_root $DATA_ROOT \
|
||||||
--seed 42 \
|
--seed 42 \
|
||||||
--mixed_precision "bf16" \
|
|
||||||
--output_dir $OUTPUT_PATH \
|
--output_dir $OUTPUT_PATH \
|
||||||
--train_batch_size 1 \
|
--train_batch_size 1 \
|
||||||
--dataloader_num_workers 4 \
|
--dataloader_num_workers 4 \
|
||||||
--pin_memory \
|
--pin_memory \
|
||||||
--caption_dropout 0.1 \
|
--caption_dropout 0.1 \
|
||||||
--max_train_steps 2000 \
|
--max_train_steps 2000 \
|
||||||
--checkpointing_steps 200 \
|
|
||||||
--checkpoints_total_limit 1 \
|
|
||||||
--gradient_accumulation_steps 4 \
|
|
||||||
--gradient_checkpointing \
|
--gradient_checkpointing \
|
||||||
--enable_slicing \
|
--enable_slicing \
|
||||||
--enable_tiling \
|
--enable_tiling \
|
||||||
--enable_model_cpu_offload \
|
--enable_model_cpu_offload \
|
||||||
--optimizer adamw --use_8bit \
|
--optimizer adamw \
|
||||||
--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 \"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 ::: \
|
--validation_prompt_separator ::: \
|
||||||
--num_validation_videos 1 \
|
--num_validation_videos 1 \
|
||||||
--validation_epochs 1 \
|
--validation_epochs 1 \
|
||||||
--allow_tf32 \
|
--allow_tf32 \
|
||||||
--report_to wandb \
|
--report_to wandb \
|
||||||
--nccl_timeout 1800"
|
--push_to_hub"
|
||||||
|
|
||||||
echo "Running command: $cmd"
|
echo "Running command: $cmd"
|
||||||
eval $cmd
|
eval $cmd
|
||||||
|
|||||||
Reference in New Issue
Block a user