LTX Video (#123)

* rename files

* ltx finetuning

* update

* update

* improvements

* make style

* gradient clipping

* update

* fix distributed inference

* update

* update
This commit is contained in:
Aryan
2024-12-19 04:27:36 +05:30
committed by GitHub
parent 80d1150a0e
commit 9ef58e2f3a
36 changed files with 3089 additions and 154 deletions
+2
View File
@@ -168,5 +168,7 @@ cython_debug/
wandb/
*.txt
dump*
outputs*
*.slurm
!requirements.txt
+103 -150
View File
@@ -1,8 +1,8 @@
# CogVideoX Factory 🧪
# finetrainers 🧪
[中文阅读](./README_zh.md)
`cogvideox-factory` was renamed to `finetrainers`. If you're looking to train CogVideoX or Mochi with the legacy training scripts, please refer to [this](./training/README.md) README instead. Everything in the `training/` directory will be eventually moved and supported under `finetrainers`.
Fine-tune Cog family of video models for custom video generation under 24GB of GPU memory ⚡️📼
FineTrainers is a work-in-progress library to support training of video models. The first priority is to support lora training for all models in [Diffusers](https://github.com/huggingface/diffusers), and eventually other methods like controlnets, control-loras, distillation, etc.
<table align="center">
<tr>
@@ -10,8 +10,6 @@ Fine-tune Cog family of video models for custom video generation under 24GB of G
</tr>
</table>
**Update 29 Nov 2024**: We have added an experimental memory-efficient trainer for Mochi-1. Check it out [here](https://github.com/a-r-r-o-w/cogvideox-factory/blob/main/training/mochi-1/)!
## Quickstart
Clone the repository and make sure the requirements are installed: `pip install -r requirements.txt` and install diffusers from source by `pip install git+https://github.com/huggingface/diffusers`.
@@ -25,152 +23,125 @@ huggingface-cli download \
--local-dir video-dataset-disney
```
Then launch LoRA fine-tuning for text-to-video (modify the different hyperparameters, dataset root, and other configuration options as per your choice):
Then launch LoRA fine-tuning. For CogVideoX and Mochi, refer to [this](./training/README.md) and [this](./training/mochi-1/README.md).
<details>
<summary> LTX Video </summary>
### Training:
```bash
# For LoRA finetuning of the text-to-video CogVideoX models
./train_text_to_video_lora.sh
#!/bin/bash
# For full finetuning of the text-to-video CogVideoX models
./train_text_to_video_sft.sh
# export TORCH_LOGS="+dynamo,recompiles,graph_breaks"
# export TORCHDYNAMO_VERBOSE=1
export WANDB_MODE="offline"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export FINETRAINERS_LOG_LEVEL=DEBUG
# For LoRA finetuning of the image-to-video CogVideoX models
./train_image_to_video_lora.sh
# Modify this based on the number of GPUs available
GPU_IDS="0,1"
DATA_ROOT="/path/to/dataset/cakify"
CAPTION_COLUMN="prompts.txt"
VIDEO_COLUMN="videos.txt"
OUTPUT_DIR="/path/to/output/directory/ltx-video/ltxv_cakify"
# Model arguments
model_cmd="--model_name ltx_video \
--pretrained_model_name_or_path Lightricks/LTX-Video"
# Dataset arguments
dataset_cmd="--data_root $DATA_ROOT \
--video_column $VIDEO_COLUMN \
--caption_column $CAPTION_COLUMN \
--id_token BW_STYLE \
--video_resolution_buckets 17x512x768 49x512x768 61x512x768 129x512x768 \
--caption_dropout_p 0.05"
# Dataloader arguments
dataloader_cmd="--dataloader_num_workers 0"
# Diffusion arguments
diffusion_cmd="--flow_resolution_shifting"
# Training arguments
training_cmd="--training_type lora \
--seed 42 \
--mixed_precision bf16 \
--batch_size 1 \
--train_steps 2000 \
--rank 128 \
--lora_alpha 128 \
--target_modules to_q to_k to_v to_out.0 \
--gradient_accumulation_steps 1 \
--gradient_checkpointing \
--checkpointing_steps 500 \
--checkpointing_limit 2 \
--enable_slicing \
--enable_tiling"
# Optimizer arguments
optimizer_cmd="--optimizer adamw \
--lr 1e-5 \
--lr_scheduler constant \
--lr_warmup_steps 100 \
--lr_num_cycles 1 \
--beta1 0.9 \
--beta2 0.95 \
--weight_decay 1e-4 \
--epsilon 1e-8 \
--max_grad_norm 1.0"
# Validation arguments
validation_cmd="--validation_prompts \"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@@@49x512x768:::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@@@129x512x768:::BW_STYLE A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance@@@49x512x768\" \
--num_validation_videos 1 \
--validation_steps 100"
# Miscellaneous arguments
miscellaneous_cmd="--tracker_name finetrainers-ltxv \
--output_dir $OUTPUT_DIR \
--nccl_timeout 1800 \
--report_to wandb"
cmd="accelerate launch --config_file accelerate_configs/uncompiled_2.yaml --gpu_ids $GPU_IDS train.py \
$model_cmd \
$dataset_cmd \
$dataloader_cmd \
$diffusion_cmd \
$training_cmd \
$optimizer_cmd \
$validation_cmd \
$miscellaneous_cmd"
echo "Running command: $cmd"
eval $cmd
echo -ne "-------------------- Finished executing script --------------------\n\n"
```
### Inference:
Assuming your LoRA is saved and pushed to the HF Hub, and named `my-awesome-name/my-awesome-lora`, we can now use the finetuned model for inference:
```diff
import torch
from diffusers import CogVideoXPipeline
from diffusers import LTXPipeline
from diffusers.utils import export_to_video
pipe = CogVideoXPipeline.from_pretrained(
"THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16
pipe = LTXPipeline.from_pretrained(
"Lightricks/LTX-Video", torch_dtype=torch.bfloat16
).to("cuda")
+ pipe.load_lora_weights("my-awesome-name/my-awesome-lora", adapter_name="cogvideox-lora")
+ pipe.set_adapters(["cogvideox-lora"], [1.0])
+ pipe.load_lora_weights("my-awesome-name/my-awesome-lora", adapter_name="ltxv-lora")
+ pipe.set_adapters(["ltxv-lora"], [1.0])
video = pipe("<my-awesome-prompt>").frames[0]
export_to_video(video, "output.mp4", fps=8)
```
For Image-to-Video LoRAs trained with multiresolution videos, one must also add the following lines (see [this](https://github.com/a-r-r-o-w/cogvideox-factory/issues/26) Issue for more details):
</details>
```python
from diffusers import CogVideoXImageToVideoPipeline
pipe = CogVideoXImageToVideoPipeline.from_pretrained(
"THUDM/CogVideoX-5b-I2V", torch_dtype=torch.bfloat16
).to("cuda")
# ...
del pipe.transformer.patch_embed.pos_embedding
pipe.transformer.patch_embed.use_learned_positional_embeddings = False
pipe.transformer.config.use_learned_positional_embeddings = False
```
You can also check if your LoRA is correctly mounted [here](tests/test_lora_inference.py).
Below we provide additional sections detailing on more options explored in this repository. They all attempt to make fine-tuning for video models as accessible as possible by reducing memory requirements as much as possible.
## Prepare Dataset and Training
Before starting the training, please check whether the dataset has been prepared according to the [dataset specifications](assets/dataset.md). We provide training scripts suitable for text-to-video and image-to-video generation, compatible with the [CogVideoX model family](https://huggingface.co/collections/THUDM/cogvideo-66c08e62f1685a3ade464cce). Training can be started using the `train*.sh` scripts, depending on the task you want to train. Let's take LoRA fine-tuning for text-to-video as an example.
- Configure environment variables as per your choice:
```bash
export TORCH_LOGS="+dynamo,recompiles,graph_breaks"
export TORCHDYNAMO_VERBOSE=1
export WANDB_MODE="offline"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
```
- Configure which GPUs to use for training: `GPU_IDS="0,1"`
- Choose hyperparameters for training. Let's try to do a sweep on learning rate and optimizer type as an example:
```bash
LEARNING_RATES=("1e-4" "1e-3")
LR_SCHEDULES=("cosine_with_restarts")
OPTIMIZERS=("adamw" "adam")
MAX_TRAIN_STEPS=("3000")
```
- Select which Accelerate configuration you would like to train with: `ACCELERATE_CONFIG_FILE="accelerate_configs/uncompiled_1.yaml"`. We provide some default configurations in the `accelerate_configs/` directory - single GPU uncompiled/compiled, 2x GPU DDP, DeepSpeed, etc. You can create your own config files with custom settings using `accelerate config --config_file my_config.yaml`.
- Specify the absolute paths and columns/files for captions and videos.
```bash
DATA_ROOT="/path/to/my/datasets/video-dataset-disney"
CAPTION_COLUMN="prompt.txt"
VIDEO_COLUMN="videos.txt"
```
- Launch experiments sweeping different hyperparameters:
```
for learning_rate in "${LEARNING_RATES[@]}"; do
for lr_schedule in "${LR_SCHEDULES[@]}"; do
for optimizer in "${OPTIMIZERS[@]}"; do
for steps in "${MAX_TRAIN_STEPS[@]}"; do
output_dir="/path/to/my/models/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox_text_to_video_lora.py \
--pretrained_model_name_or_path THUDM/CogVideoX-5b \
--data_root $DATA_ROOT \
--caption_column $CAPTION_COLUMN \
--video_column $VIDEO_COLUMN \
--id_token BW_STYLE \
--height_buckets 480 \
--width_buckets 720 \
--frame_buckets 49 \
--dataloader_num_workers 8 \
--pin_memory \
--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:::BW_STYLE A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance\" \
--validation_prompt_separator ::: \
--num_validation_videos 1 \
--validation_epochs 10 \
--seed 42 \
--rank 128 \
--lora_alpha 128 \
--mixed_precision bf16 \
--output_dir $output_dir \
--max_num_frames 49 \
--train_batch_size 1 \
--max_train_steps $steps \
--checkpointing_steps 1000 \
--gradient_accumulation_steps 1 \
--gradient_checkpointing \
--learning_rate $learning_rate \
--lr_scheduler $lr_schedule \
--lr_warmup_steps 400 \
--lr_num_cycles 1 \
--enable_slicing \
--enable_tiling \
--optimizer $optimizer \
--beta1 0.9 \
--beta2 0.95 \
--weight_decay 0.001 \
--max_grad_norm 1.0 \
--allow_tf32 \
--report_to wandb \
--nccl_timeout 1800"
echo "Running command: $cmd"
eval $cmd
echo -ne "-------------------- Finished executing script --------------------\n\n"
done
done
done
done
```
To understand what the different parameters mean, you could either take a look at the [args](./training/args.py) file or run the training script with `--help`.
Note: Training scripts are untested on MPS, so performance and memory requirements can differ widely compared to the CUDA reports below.
If you would like to use a custom dataset, refer to the dataset preparation guide [here](./assets/dataset.md).
## Memory requirements
@@ -440,21 +411,3 @@ With `train_batch_size = 4`:
<td align="center"><img src="assets/slaying-ooms.png" style="width: 480px; height: 480px;"></td>
</tr>
</table>
## TODOs
- [x] Make scripts compatible with DDP
- [ ] Make scripts compatible with FSDP
- [x] Make scripts compatible with DeepSpeed
- [ ] vLLM-powered captioning script
- [x] Multi-resolution/frame support in `prepare_dataset.py`
- [ ] Analyzing traces for potential speedups and removing as many syncs as possible
- [ ] Support for QLoRA (priority), and other types of high usage LoRAs methods
- [x] Test scripts with memory-efficient optimizer from bitsandbytes
- [x] Test scripts with CPUOffloadOptimizer, etc.
- [ ] Test scripts with torchao quantization, and low bit memory optimizers (Currently errors with AdamW (8/4-bit torchao))
- [ ] Test scripts with AdamW (8-bit bitsandbytes) + CPUOffloadOptimizer (with gradient offloading) (Currently errors out)
- [ ] [Sage Attention](https://github.com/thu-ml/SageAttention) (work with the authors to support backward pass, and optimize for A100)
> [!IMPORTANT]
> Since our goal is to make the scripts as memory-friendly as possible we don't guarantee multi-GPU training.
+17
View File
@@ -0,0 +1,17 @@
compute_environment: LOCAL_MACHINE
debug: false
distributed_type: MULTI_GPU
downcast_bf16: 'no'
enable_cpu_affinity: false
gpu_ids: all
machine_rank: 0
main_training_function: main
mixed_precision: bf16
num_machines: 1
num_processes: 8
rdzv_backend: static
same_network: true
tpu_env: []
tpu_use_cluster: false
tpu_use_sudo: false
use_cpu: false
+2
View File
@@ -0,0 +1,2 @@
from .args import Args, parse_arguments
from .trainer import Trainer
+779
View File
@@ -0,0 +1,779 @@
import argparse
from typing import Any, Dict, List, Optional, Tuple
from .constants import DEFAULT_IMAGE_RESOLUTION_BUCKETS, DEFAULT_VIDEO_RESOLUTION_BUCKETS
class Args:
r"""
The arguments for the finetrainers training script.
Args:
flow_resolution_shifting (`bool`, defaults to `False`):
Resolution-dependant shifting of timestep schedules.
[Scaling Rectified Flow Transformers for High-Resolution Image Synthesis](https://arxiv.org/abs/2403.03206)
"""
# Model arguments
model_name: str = None
pretrained_model_name_or_path: str = None
revision: Optional[str] = None
variant: Optional[str] = None
cache_dir: Optional[str] = None
# Dataset arguments
data_root: str = None
dataset_file: Optional[str] = None
video_column: str = None
caption_column: str = None
id_token: Optional[str] = None
image_resolution_buckets: List[Tuple[int, int]] = None
video_resolution_buckets: List[Tuple[int, int, int]] = None
video_reshape_mode: Optional[str] = None
caption_dropout_p: float = 0.00
caption_dropout_technique: str = "empty"
# Dataloader arguments
dataloader_num_workers: int = 0
pin_memory: bool = False
# Diffusion arguments
flow_resolution_shifting: bool = False
flow_base_image_seq_len: int = 256
flow_max_image_seq_len: int = 4096
flow_base_shift: float = 0.5
flow_max_shift: float = 1.15
flow_shift: float = 1.0
flow_weighting_scheme: str = "none"
flow_logit_mean: float = 0.0
flow_logit_std: float = 1.0
flow_mode_scale: float = 1.29
# Training arguments
training_type: str = None
seed: int = 42
mixed_precision: str = None
batch_size: int = 1
train_epochs: int = 1
train_steps: int = None
rank: int = 128
lora_alpha: float = 64
target_modules: List[str] = ["to_k", "to_q", "to_v", "to_out.0"]
gradient_accumulation_steps: int = 1
gradient_checkpointing: bool = False
checkpointing_steps: int = 500
checkpointing_limit: Optional[int] = None
resume_from_checkpoint: Optional[str] = None
enable_slicing: bool = False
enable_tiling: bool = False
# Optimizer arguments
optimizer: str = "adamw"
lr: float = 1e-4
scale_lr: bool = False
lr_scheduler: str = "cosine_with_restarts"
lr_warmup_steps: int = 0
lr_num_cycles: int = 1
lr_power: float = 1.0
beta1: float = 0.9
beta2: float = 0.95
beta3: float = 0.999
weight_decay: float = 0.0001
epsilon: float = 1e-8
max_grad_norm: float = 1.0
# Validation arguments
validation_prompts: List[str] = None
validation_images: List[str] = None
validation_videos: List[str] = None
validation_heights: List[int] = None
validation_widths: List[int] = None
validation_num_frames: List[int] = None
num_validation_videos_per_prompt: int = 1
validation_every_n_epochs: Optional[int] = None
validation_every_n_steps: Optional[int] = None
enable_model_cpu_offload: bool = False
# Miscellaneous arguments
tracker_name: str = "finetrainers"
push_to_hub: bool = False
hub_token: Optional[str] = None
hub_model_id: Optional[str] = None
output_dir: str = None
logging_dir: Optional[str] = "logs"
allow_tf32: bool = False
nccl_timeout: int = 1800 # 30 minutes
report_to: str = "wandb"
def to_dict(self) -> Dict[str, Any]:
return {
"model_arguments": {
"model_name": self.model_name,
"pretrained_model_name_or_path": self.pretrained_model_name_or_path,
"revision": self.revision,
"variant": self.variant,
"cache_dir": self.cache_dir,
},
"dataset_arguments": {
"data_root": self.data_root,
"dataset_file": self.dataset_file,
"video_column": self.video_column,
"caption_column": self.caption_column,
"id_token": self.id_token,
"image_resolution_buckets": self.image_resolution_buckets,
"video_resolution_buckets": self.video_resolution_buckets,
"video_reshape_mode": self.video_reshape_mode,
"caption_dropout_p": self.caption_dropout_p,
},
"dataloader_arguments": {
"dataloader_num_workers": self.dataloader_num_workers,
"pin_memory": self.pin_memory,
},
"training_arguments": {
"training_type": self.training_type,
"seed": self.seed,
"mixed_precision": self.mixed_precision,
"batch_size": self.batch_size,
"train_epochs": self.train_epochs,
"train_steps": self.train_steps,
"rank": self.rank,
"lora_alpha": self.lora_alpha,
"target_modules": self.target_modules,
"gradient_accumulation_steps": self.gradient_accumulation_steps,
"gradient_checkpointing": self.gradient_checkpointing,
"checkpointing_steps": self.checkpointing_steps,
"checkpointing_limit": self.checkpointing_limit,
"resume_from_checkpoint": self.resume_from_checkpoint,
"enable_slicing": self.enable_slicing,
"enable_tiling": self.enable_tiling,
},
"optimizer_arguments": {
"optimizer": self.optimizer,
"lr": self.lr,
"scale_lr": self.scale_lr,
"lr_scheduler": self.lr_scheduler,
"lr_warmup_steps": self.lr_warmup_steps,
"lr_num_cycles": self.lr_num_cycles,
"lr_power": self.lr_power,
"beta1": self.beta1,
"beta2": self.beta2,
"beta3": self.beta3,
"weight_decay": self.weight_decay,
"epsilon": self.epsilon,
"max_grad_norm": self.max_grad_norm,
},
"validation_arguments": {
"validation_prompts": self.validation_prompts,
"validation_images": self.validation_images,
"validation_videos": self.validation_videos,
"num_validation_videos_per_prompt": self.num_validation_videos_per_prompt,
"validation_every_n_epochs": self.validation_every_n_epochs,
"validation_every_n_steps": self.validation_every_n_steps,
"enable_model_cpu_offload": self.enable_model_cpu_offload,
},
"miscellaneous_arguments": {
"tracker_name": self.tracker_name,
"push_to_hub": self.push_to_hub,
"hub_token": self.hub_token,
"hub_model_id": self.hub_model_id,
"output_dir": self.output_dir,
"logging_dir": self.logging_dir,
"allow_tf32": self.allow_tf32,
"nccl_timeout": self.nccl_timeout,
"report_to": self.report_to,
},
}
def parse_arguments() -> Args:
parser = argparse.ArgumentParser()
_add_model_arguments(parser)
_add_dataset_arguments(parser)
_add_dataloader_arguments(parser)
_add_diffusion_arguments(parser)
_add_training_arguments(parser)
_add_optimizer_arguments(parser)
_add_validation_arguments(parser)
_add_miscellaneous_arguments(parser)
args = parser.parse_args()
return _map_to_args_type(args)
def validate_args(args: Args):
_validate_validation_args(args)
def _add_model_arguments(parser: argparse.ArgumentParser) -> None:
parser.add_argument("--model_name", type=str, required=True, choices=["ltx_video"], help="Name of model to train.")
parser.add_argument(
"--pretrained_model_name_or_path",
type=str,
default=None,
help="Path to pretrained model or model identifier from huggingface.co/models.",
)
parser.add_argument(
"--revision",
type=str,
default=None,
required=False,
help="Revision of pretrained model identifier from huggingface.co/models.",
)
parser.add_argument(
"--variant",
type=str,
default=None,
help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16",
)
parser.add_argument(
"--cache_dir",
type=str,
default=None,
help="The directory where the downloaded models and datasets will be stored.",
)
def _add_dataset_arguments(parser: argparse.ArgumentParser) -> None:
def parse_resolution_bucket(resolution_bucket: str) -> Tuple[int, ...]:
return tuple(map(int, resolution_bucket.split("x")))
def parse_image_resolution_bucket(resolution_bucket: str) -> Tuple[int, int]:
resolution_bucket = parse_resolution_bucket(resolution_bucket)
assert (
len(resolution_bucket) == 2
), f"Expected 2D resolution bucket, got {len(resolution_bucket)}D resolution bucket"
return resolution_bucket
def parse_video_resolution_bucket(resolution_bucket: str) -> Tuple[int, int, int]:
resolution_bucket = parse_resolution_bucket(resolution_bucket)
assert (
len(resolution_bucket) == 3
), f"Expected 3D resolution bucket, got {len(resolution_bucket)}D resolution bucket"
return resolution_bucket
parser.add_argument(
"--data_root",
type=str,
default=None,
help=("A folder containing the training data."),
)
parser.add_argument(
"--dataset_file",
type=str,
default=None,
help=("Path to a CSV file if loading prompts/video paths using this format."),
)
parser.add_argument(
"--video_column",
type=str,
default="video",
help="The column of the dataset containing videos. Or, the name of the file in `--data_root` folder containing the line-separated path to video data.",
)
parser.add_argument(
"--caption_column",
type=str,
default="text",
help="The column of the dataset containing the instance prompt for each video. Or, the name of the file in `--data_root` folder containing the line-separated instance prompts.",
)
parser.add_argument(
"--id_token",
type=str,
default=None,
help="Identifier token appended to the start of each prompt if provided.",
)
parser.add_argument(
"--image_resolution_buckets",
type=parse_image_resolution_bucket,
default=None,
nargs="+",
help="Resolution buckets for images.",
)
parser.add_argument(
"--video_resolution_buckets",
type=parse_video_resolution_bucket,
default=None,
nargs="+",
help="Resolution buckets for videos.",
)
parser.add_argument(
"--video_reshape_mode",
type=str,
default=None,
help="All input videos are reshaped to this mode. Choose between ['center', 'random', 'none']",
)
parser.add_argument(
"--caption_dropout_p",
type=float,
default=0.00,
help="Probability of dropout for the caption tokens.",
)
parser.add_argument(
"--caption_dropout_technique",
type=str,
default="empty",
choices=["empty", "zero"],
help="Technique to use for caption dropout.",
)
def _add_dataloader_arguments(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--dataloader_num_workers",
type=int,
default=0,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--pin_memory",
action="store_true",
help="Whether or not to use the pinned memory setting in pytorch dataloader.",
)
def _add_diffusion_arguments(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--flow_resolution_shifting",
action="store_true",
help="Resolution-dependant shifting of timestep schedules.",
)
parser.add_argument(
"--flow_weighting_scheme",
type=str,
default="none",
choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"],
help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'),
)
parser.add_argument(
"--flow_logit_mean",
type=float,
default=0.0,
help="mean to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--flow_logit_std",
type=float,
default=1.0,
help="std to use when using the `'logit_normal'` weighting scheme.",
)
parser.add_argument(
"--flow_mode_scale",
type=float,
default=1.29,
help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.",
)
def _add_training_arguments(parser: argparse.ArgumentParser) -> None:
# TODO: support full finetuning and other kinds
parser.add_argument(
"--training_type",
type=str,
default=None,
help="Type of training to perform. Choose between ['lora']",
)
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
parser.add_argument(
"--mixed_precision",
type=str,
default="no",
choices=["no", "fp8", "fp16", "bf16"],
help=(
"Whether to use mixed precision. Defaults 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(
"--batch_size",
type=int,
default=4,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--train_epochs", type=int, default=1)
parser.add_argument(
"--train_steps",
type=int,
default=None,
help="Total number of training steps to perform. If provided, overrides `--num_train_epochs`.",
)
parser.add_argument("--rank", type=int, default=64, help="The rank for LoRA matrices.")
parser.add_argument(
"--lora_alpha",
type=int,
default=64,
help="The lora_alpha to compute scaling factor (lora_alpha / rank) for LoRA matrices.",
)
parser.add_argument(
"--target_modules", type=str, default="to_k to_q to_v to_out.0", nargs="+", help="The target modules for LoRA."
)
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",
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
)
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(
"--checkpointing_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(
"--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 _add_optimizer_arguments(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--lr",
type=float,
default=1e-4,
help="Initial learning rate (after the potential warmup period) to use.",
)
parser.add_argument(
"--scale_lr",
action="store_true",
default=False,
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
)
parser.add_argument(
"--lr_scheduler",
type=str,
default="constant",
help=(
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
' "constant", "constant_with_warmup"]'
),
)
parser.add_argument(
"--lr_warmup_steps",
type=int,
default=500,
help="Number of steps for the warmup in the lr scheduler.",
)
parser.add_argument(
"--lr_num_cycles",
type=int,
default=1,
help="Number of hard resets of the lr in cosine_with_restarts scheduler.",
)
parser.add_argument(
"--lr_power",
type=float,
default=1.0,
help="Power factor of the polynomial scheduler.",
)
parser.add_argument(
"--optimizer",
type=lambda s: s.lower(),
default="adam",
choices=["adam", "adamw"],
help=("The optimizer type to use."),
)
parser.add_argument(
"--beta1",
type=float,
default=0.9,
help="The beta1 parameter for the Adam and Prodigy optimizers.",
)
parser.add_argument(
"--beta2",
type=float,
default=0.95,
help="The beta2 parameter for the Adam and Prodigy optimizers.",
)
parser.add_argument(
"--beta3",
type=float,
default=None,
help="Coefficients for computing the Prodigy optimizer's stepsize using running averages. If set to None, uses the value of square root of beta2.",
)
parser.add_argument(
"--weight_decay",
type=float,
default=1e-04,
help="Weight decay to use for optimizer.",
)
parser.add_argument(
"--epsilon",
type=float,
default=1e-8,
help="Epsilon value for the Adam optimizer and Prodigy optimizers.",
)
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
def _add_validation_arguments(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--validation_prompts",
type=str,
default=None,
help="One or more prompt(s) that is used during validation to verify that the model is learning. Multiple validation prompts should be separated by the '--validation_prompt_seperator' string.",
)
parser.add_argument(
"--validation_images",
type=str,
default=None,
help="One or more image path(s)/URLs that is used during validation to verify that the model is learning. Multiple validation paths should be separated by the '--validation_prompt_seperator' string. These should correspond to the order of the validation prompts.",
)
parser.add_argument(
"--validation_videos",
type=str,
default=None,
help="One or more video path(s)/URLs that is used during validation to verify that the model is learning. Multiple validation paths should be separated by the '--validation_prompt_seperator' string. These should correspond to the order of the validation prompts.",
)
parser.add_argument(
"--validation_separator",
type=str,
default=":::",
help="String that separates multiple validation prompts",
)
parser.add_argument(
"--num_validation_videos",
type=int,
default=1,
help="Number of videos that should be generated during validation per `validation_prompt`.",
)
parser.add_argument(
"--validation_epochs",
type=int,
default=None,
help="Run validation every X training epochs. Validation consists of running the validation prompt `args.num_validation_videos` times.",
)
parser.add_argument(
"--validation_steps",
type=int,
default=None,
help="Run validation every X training steps. Validation consists of running the validation prompt `args.num_validation_videos` times.",
)
parser.add_argument(
"--enable_model_cpu_offload",
action="store_true",
default=False,
help="Whether or not to enable model-wise CPU offloading when performing validation/testing to save memory.",
)
def _add_miscellaneous_arguments(parser: argparse.ArgumentParser) -> None:
parser.add_argument("--tracker_name", type=str, default="finetrainers", help="Project tracker name")
parser.add_argument(
"--push_to_hub",
action="store_true",
help="Whether or not to push the model to the Hub.",
)
parser.add_argument(
"--hub_token",
type=str,
default=None,
help="The token to use to push to the Model Hub.",
)
parser.add_argument(
"--hub_model_id",
type=str,
default=None,
help="The name of the repository to keep in sync with the local `output_dir`.",
)
parser.add_argument(
"--output_dir",
type=str,
default="finetrainer-training",
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help="Directory where logs are stored.",
)
parser.add_argument(
"--allow_tf32",
action="store_true",
help=(
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
)
parser.add_argument(
"--nccl_timeout",
type=int,
default=600,
help="Maximum timeout duration before which allgather, or related, operations fail in multi-GPU/multi-node training settings.",
)
parser.add_argument(
"--report_to",
type=str,
default="none",
help=(
'The integration to report the results and logs to. Supported platforms are `"tensorboard"`'
' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.'
),
)
def _map_to_args_type(args: Dict[str, Any]) -> Args:
result_args = Args()
# Model arguments
result_args.model_name = args.model_name
result_args.pretrained_model_name_or_path = args.pretrained_model_name_or_path
result_args.revision = args.revision
result_args.variant = args.variant
result_args.cache_dir = args.cache_dir
# Dataset arguments
if args.data_root is None and args.dataset_file is None:
raise ValueError("At least one of `data_root` or `dataset_file` should be provided.")
result_args.data_root = args.data_root
result_args.dataset_file = args.dataset_file
result_args.video_column = args.video_column
result_args.caption_column = args.caption_column
result_args.id_token = args.id_token
result_args.image_resolution_buckets = args.image_resolution_buckets or DEFAULT_IMAGE_RESOLUTION_BUCKETS
result_args.video_resolution_buckets = args.video_resolution_buckets or DEFAULT_VIDEO_RESOLUTION_BUCKETS
result_args.video_reshape_mode = args.video_reshape_mode
result_args.caption_dropout_p = args.caption_dropout_p
# Dataloader arguments
result_args.dataloader_num_workers = args.dataloader_num_workers
result_args.pin_memory = args.pin_memory
# Diffusion arguments
result_args.flow_resolution_shifting = args.flow_resolution_shifting
result_args.flow_weighting_scheme = args.flow_weighting_scheme
result_args.flow_logit_mean = args.flow_logit_mean
result_args.flow_logit_std = args.flow_logit_std
result_args.flow_mode_scale = args.flow_mode_scale
# Training arguments
result_args.training_type = args.training_type
result_args.seed = args.seed
result_args.mixed_precision = args.mixed_precision
result_args.batch_size = args.batch_size
result_args.train_epochs = args.train_epochs
result_args.train_steps = args.train_steps
result_args.rank = args.rank
result_args.lora_alpha = args.lora_alpha
result_args.gradient_accumulation_steps = args.gradient_accumulation_steps
result_args.gradient_checkpointing = args.gradient_checkpointing
result_args.checkpointing_steps = args.checkpointing_steps
result_args.checkpointing_limit = args.checkpointing_limit
result_args.resume_from_checkpoint = args.resume_from_checkpoint
result_args.enable_slicing = args.enable_slicing
result_args.enable_tiling = args.enable_tiling
# Optimizer arguments
result_args.optimizer = args.optimizer or "adamw"
result_args.lr = args.lr or 1e-4
result_args.scale_lr = args.scale_lr
result_args.lr_scheduler = args.lr_scheduler
result_args.lr_warmup_steps = args.lr_warmup_steps
result_args.lr_num_cycles = args.lr_num_cycles
result_args.lr_power = args.lr_power
result_args.beta1 = args.beta1
result_args.beta2 = args.beta2
result_args.beta3 = args.beta3
result_args.weight_decay = args.weight_decay
result_args.epsilon = args.epsilon
result_args.max_grad_norm = args.max_grad_norm
# Validation arguments
validation_prompts = args.validation_prompts.split(args.validation_separator) if args.validation_prompts else []
validation_images = args.validation_images.split(args.validation_separator) if args.validation_images else None
validation_videos = args.validation_videos.split(args.validation_separator) if args.validation_videos else None
stripped_validation_prompts = []
validation_heights = []
validation_widths = []
validation_num_frames = []
for prompt in validation_prompts:
prompt: str
prompt = prompt.strip()
actual_prompt, separator, resolution = prompt.rpartition("@@@")
stripped_validation_prompts.append(actual_prompt)
num_frames, height, width = None, None, None
if len(resolution) > 0:
num_frames, height, width = map(int, resolution.split("x"))
validation_num_frames.append(num_frames)
validation_heights.append(height)
validation_widths.append(width)
if validation_images is None:
validation_images = [None] * len(validation_prompts)
if validation_videos is None:
validation_videos = [None] * len(validation_prompts)
result_args.validation_prompts = stripped_validation_prompts
result_args.validation_heights = validation_heights
result_args.validation_widths = validation_widths
result_args.validation_num_frames = validation_num_frames
result_args.validation_images = validation_images
result_args.validation_videos = validation_videos
result_args.num_validation_videos_per_prompt = args.num_validation_videos
result_args.validation_every_n_epochs = args.validation_epochs
result_args.validation_every_n_steps = args.validation_steps
result_args.enable_model_cpu_offload = args.enable_model_cpu_offload
# Miscellaneous arguments
result_args.tracker_name = args.tracker_name
result_args.push_to_hub = args.push_to_hub
result_args.hub_token = args.hub_token
result_args.hub_model_id = args.hub_model_id
result_args.output_dir = args.output_dir
result_args.logging_dir = args.logging_dir
result_args.allow_tf32 = args.allow_tf32
result_args.nccl_timeout = args.nccl_timeout
result_args.report_to = args.report_to
return result_args
def _validate_validation_args(args: Args):
assert args.validation_prompts is not None, "Validation prompts are required for validation"
if args.validation_images is not None:
assert len(args.validation_images) == len(
args.validation_prompts
), "Validation images and prompts should be of same length"
if args.validation_videos is not None:
assert len(args.validation_videos) == len(
args.validation_prompts
), "Validation videos and prompts should be of same length"
assert len(args.validation_prompts) == len(
args.validation_heights
), "Validation prompts and heights should be of same length"
assert len(args.validation_prompts) == len(
args.validation_widths
), "Validation prompts and widths should be of same length"
+50
View File
@@ -0,0 +1,50 @@
import os
DEFAULT_HEIGHT_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536]
DEFAULT_WIDTH_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536]
DEFAULT_FRAME_BUCKETS = [49]
DEFAULT_IMAGE_RESOLUTION_BUCKETS = []
for height in DEFAULT_HEIGHT_BUCKETS:
for width in DEFAULT_WIDTH_BUCKETS:
DEFAULT_IMAGE_RESOLUTION_BUCKETS.append((height, width))
DEFAULT_VIDEO_RESOLUTION_BUCKETS = []
for frames in DEFAULT_FRAME_BUCKETS:
for height in DEFAULT_HEIGHT_BUCKETS:
for width in DEFAULT_WIDTH_BUCKETS:
DEFAULT_VIDEO_RESOLUTION_BUCKETS.append((frames, height, width))
FINETRAINERS_LOG_LEVEL = os.environ.get("FINETRAINERS_LOG_LEVEL", "INFO")
MODEL_DESCRIPTION = r"""
\# {model_id} {training_type} finetune
<Gallery />
\#\# Model Description
This model is a {training_type} of the `{model_id}` model.
This model was trained using the `fine-video-trainers` library - a repository containing memory-optimized scripts for training video models with [Diffusers](https://github.com/huggingface/diffusers).
\#\# Download model
[Download LoRA]({repo_id}/tree/main) in the Files & Versions tab.
\#\# Usage
Requires [🧨 Diffusers](https://github.com/huggingface/diffusers) installed.
```python
{model_example}
```
For more details, including weighting, merging and fusing LoRAs, check the [documentation](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters) on loading LoRAs in diffusers.
\#\# License
Please adhere to the license of the base model.
""".strip()
+321
View File
@@ -0,0 +1,321 @@
import random
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import pandas as pd
import torch
import torchvision.transforms as TT
from accelerate.logging import get_logger
from torch.utils.data import Dataset, Sampler
from torchvision import transforms
from torchvision.transforms import InterpolationMode
from torchvision.transforms.functional import resize
# Must import after torch because this can sometimes lead to a nasty segmentation fault, or stack smashing error
# Very few bug reports but it happens. Look in decord Github issues for more relevant information.
import decord # isort:skip
decord.bridge.set_bridge("torch")
logger = get_logger(__name__)
class VideoDataset(Dataset):
def __init__(
self,
data_root: str,
caption_column: str,
video_column: str,
resolution_buckets: List[Tuple[int, int, int]],
dataset_file: Optional[str] = None,
id_token: Optional[str] = None,
) -> None:
super().__init__()
self.data_root = Path(data_root)
self.dataset_file = dataset_file
self.caption_column = caption_column
self.video_column = video_column
self.id_token = f"{id_token.strip()} " if id_token else ""
self.resolution_buckets = resolution_buckets
# Two methods of loading data are supported.
# - Using a CSV: caption_column and video_column must be some column in the CSV. One could
# make use of other columns too, such as a motion score or aesthetic score, by modifying the
# logic in CSV processing.
# - Using two files containing line-separate captions and relative paths to videos.
# For a more detailed explanation about preparing dataset format, checkout the README.
if dataset_file is None:
(
self.prompts,
self.video_paths,
) = self._load_dataset_from_local_path()
else:
(
self.prompts,
self.video_paths,
) = self._load_dataset_from_csv()
if len(self.video_paths) != len(self.prompts):
raise ValueError(
f"Expected length of prompts and videos to be the same but found {len(self.prompts)=} and {len(self.video_paths)=}. Please ensure that the number of caption prompts and videos match in your dataset."
)
self.video_transforms = transforms.Compose(
[
transforms.Lambda(self.scale_transform),
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
]
)
@staticmethod
def scale_transform(x):
return x / 255.0
def __len__(self) -> int:
return len(self.video_paths)
def __getitem__(self, index: int) -> Dict[str, Any]:
if isinstance(index, list):
# Here, index is actually a list of data objects that we need to return.
# The BucketSampler should ideally return indices. But, in the sampler, we'd like
# to have information about num_frames, height and width. Since this is not stored
# as metadata, we need to read the video to get this information. You could read this
# information without loading the full video in memory, but we do it anyway. In order
# to not load the video twice (once to get the metadata, and once to return the loaded video
# based on sampled indices), we cache it in the BucketSampler. When the sampler is
# to yield, we yield the cache data instead of indices. So, this special check ensures
# that data is not loaded a second time. PRs are welcome for improvements.
return index
prompt = self.id_token + self.prompts[index]
video = self._preprocess_video(self.video_paths[index])
return {
"prompt": prompt,
"video": video,
"video_metadata": {
"num_frames": video.shape[0],
"height": video.shape[2],
"width": video.shape[3],
},
}
def _load_dataset_from_local_path(self) -> Tuple[List[str], List[str]]:
if not self.data_root.exists():
raise ValueError("Root folder for videos does not exist")
prompt_path = self.data_root.joinpath(self.caption_column)
video_path = self.data_root.joinpath(self.video_column)
if not prompt_path.exists() or not prompt_path.is_file():
raise ValueError(
"Expected `--caption_column` to be path to a file in `--data_root` containing line-separated text prompts."
)
if not video_path.exists() or not video_path.is_file():
raise ValueError(
"Expected `--video_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory."
)
with open(prompt_path, "r", encoding="utf-8") as file:
prompts = [line.strip() for line in file.readlines() if len(line.strip()) > 0]
with open(video_path, "r", encoding="utf-8") as file:
video_paths = [self.data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0]
if any(not path.is_file() for path in video_paths):
raise ValueError(
f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file."
)
return prompts, video_paths
def _load_dataset_from_csv(self) -> Tuple[List[str], List[str]]:
df = pd.read_csv(self.dataset_file)
prompts = df[self.caption_column].tolist()
video_paths = df[self.video_column].tolist()
video_paths = [self.data_root.joinpath(line.strip()) for line in video_paths]
if any(not path.is_file() for path in video_paths):
raise ValueError(
f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file."
)
return prompts, video_paths
def _preprocess_video(self, path: Path) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
r"""
Loads a single video, or latent and prompt embedding, based on initialization parameters.
If returning a video, returns a [F, C, H, W] video tensor, and None for the prompt embedding. Here,
F, C, H and W are the frames, channels, height and width of the input video.
"""
video_reader = decord.VideoReader(uri=path.as_posix())
video_num_frames = len(video_reader)
indices = list(range(0, video_num_frames, video_num_frames // self.max_num_frames))
frames = video_reader.get_batch(indices)
frames = frames[: self.max_num_frames].float()
frames = frames.permute(0, 3, 1, 2).contiguous()
frames = torch.stack([self.video_transforms(frame) for frame in frames], dim=0)
return frames
class VideoDatasetWithResizing(VideoDataset):
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.max_num_frames = max(self.resolution_buckets, key=lambda x: x[0])[0]
def _preprocess_video(self, path: Path) -> torch.Tensor:
video_reader = decord.VideoReader(uri=path.as_posix())
video_num_frames = len(video_reader)
nearest_frame_bucket = min(
[bucket for bucket in self.resolution_buckets if bucket[0] <= video_num_frames],
key=lambda x: abs(x[0] - min(video_num_frames, self.max_num_frames)),
default=1,
)[0]
frame_indices = list(range(0, video_num_frames, video_num_frames // nearest_frame_bucket))
frames = video_reader.get_batch(frame_indices)
frames = frames[:nearest_frame_bucket].float()
frames = frames.permute(0, 3, 1, 2).contiguous()
nearest_res = self._find_nearest_resolution(frames.shape[2], frames.shape[3])
frames_resized = torch.stack([resize(frame, nearest_res) for frame in frames], dim=0)
frames = torch.stack([self.video_transforms(frame) for frame in frames_resized], dim=0)
return frames
def _find_nearest_resolution(self, height, width):
nearest_res = min(self.resolution_buckets, key=lambda x: abs(x[1] - height) + abs(x[2] - width))
return nearest_res[1], nearest_res[2]
class VideoDatasetWithResizeAndRectangleCrop(VideoDataset):
def __init__(self, video_reshape_mode: str = "center", *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.video_reshape_mode = video_reshape_mode
self.max_num_frames = max(self.resolution_buckets, key=lambda x: x[0])[0]
def _resize_for_rectangle_crop(self, arr, image_size):
reshape_mode = self.video_reshape_mode
if arr.shape[3] / arr.shape[2] > image_size[1] / image_size[0]:
arr = resize(
arr,
size=[image_size[0], int(arr.shape[3] * image_size[0] / arr.shape[2])],
interpolation=InterpolationMode.BICUBIC,
)
else:
arr = resize(
arr,
size=[int(arr.shape[2] * image_size[1] / arr.shape[3]), image_size[1]],
interpolation=InterpolationMode.BICUBIC,
)
h, w = arr.shape[2], arr.shape[3]
arr = arr.squeeze(0)
delta_h = h - image_size[0]
delta_w = w - image_size[1]
if reshape_mode == "random" or reshape_mode == "none":
top = np.random.randint(0, delta_h + 1)
left = np.random.randint(0, delta_w + 1)
elif reshape_mode == "center":
top, left = delta_h // 2, delta_w // 2
else:
raise NotImplementedError
arr = TT.functional.crop(arr, top=top, left=left, height=image_size[0], width=image_size[1])
return arr
def _preprocess_video(self, path: Path) -> torch.Tensor:
video_reader = decord.VideoReader(uri=path.as_posix())
video_num_frames = len(video_reader)
nearest_frame_bucket = min(
[bucket for bucket in self.resolution_buckets if bucket <= video_num_frames],
key=lambda x: abs(x[0] - min(video_num_frames, self.max_num_frames)),
default=1,
)[0]
frame_indices = list(range(0, video_num_frames, video_num_frames // nearest_frame_bucket))
frames = video_reader.get_batch(frame_indices)
frames = frames[:nearest_frame_bucket].float()
frames = frames.permute(0, 3, 1, 2).contiguous()
nearest_res = self._find_nearest_resolution(frames.shape[2], frames.shape[3])
frames_resized = self._resize_for_rectangle_crop(frames, nearest_res)
frames = torch.stack([self.video_transforms(frame) for frame in frames_resized], dim=0)
return frames
def _find_nearest_resolution(self, height, width):
nearest_res = min(self.resolutions, key=lambda x: abs(x[1] - height) + abs(x[2] - width))
return nearest_res[1], nearest_res[2]
class BucketSampler(Sampler):
r"""
PyTorch Sampler that groups 3D data by height, width and frames.
Args:
data_source (`VideoDataset`):
A PyTorch dataset object that is an instance of `VideoDataset`.
batch_size (`int`, defaults to `8`):
The batch size to use for training.
shuffle (`bool`, defaults to `True`):
Whether or not to shuffle the data in each batch before dispatching to dataloader.
drop_last (`bool`, defaults to `False`):
Whether or not to drop incomplete buckets of data after completely iterating over all data
in the dataset. If set to True, only batches that have `batch_size` number of entries will
be yielded. If set to False, it is guaranteed that all data in the dataset will be processed
and batches that do not have `batch_size` number of entries will also be yielded.
"""
def __init__(
self, data_source: VideoDataset, batch_size: int = 8, shuffle: bool = True, drop_last: bool = False
) -> None:
self.data_source = data_source
self.batch_size = batch_size
self.shuffle = shuffle
self.drop_last = drop_last
self.buckets = {resolution: [] for resolution in data_source.resolution_buckets}
self._raised_warning_for_drop_last = False
def __len__(self):
if self.drop_last and not self._raised_warning_for_drop_last:
self._raised_warning_for_drop_last = True
logger.warning(
"Calculating the length for bucket sampler is not possible when `drop_last` is set to True. This may cause problems when setting the number of epochs used for training."
)
return (len(self.data_source) + self.batch_size - 1) // self.batch_size
def __iter__(self):
for index, data in enumerate(self.data_source):
video_metadata = data["video_metadata"]
f, h, w = video_metadata["num_frames"], video_metadata["height"], video_metadata["width"]
self.buckets[(f, h, w)].append(data)
if len(self.buckets[(f, h, w)]) == self.batch_size:
if self.shuffle:
random.shuffle(self.buckets[(f, h, w)])
yield self.buckets[(f, h, w)]
del self.buckets[(f, h, w)]
self.buckets[(f, h, w)] = []
if self.drop_last:
return
for fhw, bucket in list(self.buckets.items()):
if len(bucket) == 0:
continue
if self.shuffle:
random.shuffle(bucket)
yield bucket
del self.buckets[fhw]
self.buckets[fhw] = []
+1
View File
@@ -0,0 +1 @@
from .ltx_video import LTX_VIDEO_T2V_CONFIG
+264
View File
@@ -0,0 +1,264 @@
from typing import Dict, List, Optional, Union
import torch
import torch.nn as nn
from accelerate.logging import get_logger
from diffusers import AutoencoderKLLTXVideo, FlowMatchEulerDiscreteScheduler, LTXPipeline, LTXVideoTransformer3DModel
from diffusers.utils import logging
from transformers import T5EncoderModel, T5Tokenizer
from PIL import Image
logger = get_logger("finetrainers") # pylint: disable=invalid-name
def load_components(
model_id: str = "Lightricks/LTX-Video",
text_encoder_dtype: torch.dtype = torch.bfloat16,
transformer_dtype: torch.dtype = torch.bfloat16,
vae_dtype: torch.dtype = torch.bfloat16,
cache_dir: Optional[str] = None,
) -> Dict[str, nn.Module]:
tokenizer = T5Tokenizer.from_pretrained(model_id, subfolder="tokenizer", cache_dir=cache_dir)
text_encoder = T5EncoderModel.from_pretrained(
model_id, subfolder="text_encoder", torch_dtype=text_encoder_dtype, cache_dir=cache_dir
)
transformer = LTXVideoTransformer3DModel.from_pretrained(
model_id, subfolder="transformer", torch_dtype=transformer_dtype, cache_dir=cache_dir
)
vae = AutoencoderKLLTXVideo.from_pretrained(model_id, subfolder="vae", torch_dtype=vae_dtype, cache_dir=cache_dir)
scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(model_id, subfolder="scheduler", cache_dir=cache_dir)
return {
"tokenizer": tokenizer,
"text_encoder": text_encoder,
"transformer": transformer,
"vae": vae,
"scheduler": scheduler,
}
def initialize_pipeline(
model_id: str = "Lightricks/LTX-Video",
text_encoder_dtype: torch.dtype = torch.bfloat16,
transformer_dtype: torch.dtype = torch.bfloat16,
vae_dtype: torch.dtype = torch.bfloat16,
tokenizer: Optional[T5Tokenizer] = None,
text_encoder: Optional[T5EncoderModel] = None,
transformer: Optional[LTXVideoTransformer3DModel] = None,
vae: Optional[AutoencoderKLLTXVideo] = None,
scheduler: Optional[FlowMatchEulerDiscreteScheduler] = None,
device: Optional[torch.device] = None,
cache_dir: Optional[str] = None,
enable_slicing: bool = False,
enable_tiling: bool = False,
enable_model_cpu_offload: bool = False,
) -> LTXPipeline:
component_name_pairs = [
("tokenizer", tokenizer),
("text_encoder", text_encoder),
("transformer", transformer),
("vae", vae),
("scheduler", scheduler),
]
components = {}
for name, component in component_name_pairs:
if component is not None:
components[name] = component
pipe = LTXPipeline.from_pretrained(model_id, **components, cache_dir=cache_dir)
pipe.text_encoder = pipe.text_encoder.to(dtype=text_encoder_dtype)
pipe.transformer = pipe.transformer.to(dtype=transformer_dtype)
pipe.vae = pipe.vae.to(dtype=vae_dtype)
if enable_slicing:
pipe.vae.enable_slicing()
if enable_tiling:
pipe.vae.enable_tiling()
if enable_model_cpu_offload:
pipe.enable_model_cpu_offload(device=device)
else:
pipe.to(device=device)
return pipe
def prepare_conditions(
tokenizer: T5Tokenizer,
text_encoder: T5EncoderModel,
prompt: Union[str, List[str]],
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
max_sequence_length: int = 128,
) -> torch.Tensor:
device = device or text_encoder.device
dtype = dtype or text_encoder.dtype
if isinstance(prompt, str):
prompt = [prompt]
return _encode_prompt_t5(tokenizer, text_encoder, prompt, device, dtype, max_sequence_length)
def prepare_latents(
vae: AutoencoderKLLTXVideo,
image_or_video: torch.Tensor,
patch_size: int = 1,
patch_size_t: int = 1,
device: Optional[torch.device] = None,
dtype: Optional[torch.dtype] = None,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
device = device or vae.device
dtype = dtype or vae.dtype
if image_or_video.ndim == 4:
image_or_video = image_or_video.unsqueeze(2)
assert image_or_video.ndim == 5, f"Expected 5D tensor, got {image_or_video.ndim}D tensor"
image_or_video = image_or_video.to(device=device, dtype=dtype)
image_or_video = image_or_video.permute(0, 2, 1, 3, 4).contiguous() # [B, C, F, H, W] -> [B, F, C, H, W]
latents = vae.encode(image_or_video).latent_dist.sample(generator=generator)
_, _, num_frames, height, width = latents.shape
latents = _normalize_latents(latents, vae.latents_mean, vae.latents_std)
latents = _pack_latents(latents, patch_size, patch_size_t)
return {"latents": latents, "num_frames": num_frames, "height": height, "width": width}
def collate_fn_t2v(batch: List[List[Dict[str, torch.Tensor]]]) -> Dict[str, torch.Tensor]:
return {
"prompts": [x["prompt"] for x in batch[0]],
"videos": torch.stack([x["video"] for x in batch[0]]),
}
def forward_pass(
transformer: LTXVideoTransformer3DModel,
prompt_embeds: torch.Tensor,
prompt_attention_mask: torch.Tensor,
latents: torch.Tensor,
noisy_latents: torch.Tensor,
timesteps: torch.LongTensor,
num_frames: int,
height: int,
width: int,
) -> torch.Tensor:
# TODO(aryan): make configurable
rope_interpolation_scale = [1 / 25, 32, 32]
denoised_latents = transformer(
hidden_states=noisy_latents,
encoder_hidden_states=prompt_embeds,
timestep=timesteps,
encoder_attention_mask=prompt_attention_mask,
num_frames=num_frames,
height=height,
width=width,
rope_interpolation_scale=rope_interpolation_scale,
return_dict=False,
)[0]
return {"latents": denoised_latents}
def validation(
pipeline: LTXPipeline,
prompt: str,
image: Optional[Image.Image] = None,
video: Optional[List[Image.Image]] = None,
height: Optional[int] = None,
width: Optional[int] = None,
num_frames: Optional[int] = None,
frame_rate: int = 25,
num_videos_per_prompt: int = 1,
generator: Optional[torch.Generator] = None,
**kwargs,
):
generation_kwargs = {
"prompt": prompt,
"height": height,
"width": width,
"num_frames": num_frames,
"frame_rate": frame_rate,
"num_videos_per_prompt": num_videos_per_prompt,
"generator": generator,
"return_dict": True,
"output_type": "pil",
}
generation_kwargs = {k: v for k, v in generation_kwargs.items() if v is not None}
video = pipeline(**generation_kwargs).frames[0]
return [("video", video)]
def _encode_prompt_t5(
tokenizer: T5Tokenizer,
text_encoder: T5EncoderModel,
prompt: List[str],
device: torch.device,
dtype: torch.dtype,
max_sequence_length,
) -> torch.Tensor:
batch_size = len(prompt)
text_inputs = tokenizer(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
add_special_tokens=True,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.bool().to(device)
prompt_embeds = text_encoder(text_input_ids.to(device))[0]
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
return {"prompt_embeds": prompt_embeds, "prompt_attention_mask": prompt_attention_mask}
def _normalize_latents(
latents: torch.Tensor, latents_mean: torch.Tensor, latents_std: torch.Tensor, scaling_factor: float = 1.0
) -> torch.Tensor:
# Normalize latents across the channel dimension [B, C, F, H, W]
latents_mean = latents_mean.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype)
latents_std = latents_std.view(1, -1, 1, 1, 1).to(latents.device, latents.dtype)
latents = (latents - latents_mean) * scaling_factor / latents_std
return latents
def _pack_latents(latents: torch.Tensor, patch_size: int = 1, patch_size_t: int = 1) -> torch.Tensor:
# Unpacked latents of shape are [B, C, F, H, W] are patched into tokens of shape [B, C, F // p_t, p_t, H // p, p, W // p, p].
# The patch dimensions are then permuted and collapsed into the channel dimension of shape:
# [B, F // p_t * H // p * W // p, C * p_t * p * p] (an ndim=3 tensor).
# dim=0 is the batch size, dim=1 is the effective video sequence length, dim=2 is the effective number of input features
batch_size, num_channels, num_frames, height, width = latents.shape
post_patch_num_frames = num_frames // patch_size_t
post_patch_height = height // patch_size
post_patch_width = width // patch_size
latents = latents.reshape(
batch_size,
-1,
post_patch_num_frames,
patch_size_t,
post_patch_height,
patch_size,
post_patch_width,
patch_size,
)
latents = latents.permute(0, 2, 4, 6, 1, 3, 5, 7).flatten(4, 7).flatten(1, 3)
return latents
LTX_VIDEO_T2V_CONFIG = {
"pipeline_cls": LTXPipeline,
"load_components": load_components,
"initialize_pipeline": initialize_pipeline,
"prepare_conditions": prepare_conditions,
"prepare_latents": prepare_latents,
"collate_fn": collate_fn_t2v,
"forward_pass": forward_pass,
"validation": validation,
}
+16
View File
@@ -0,0 +1,16 @@
from typing import Any, Dict
from .ltx_video import LTX_VIDEO_T2V_CONFIG
SUPPORTED_MODEL_CONFIGS = {
"ltx_video": LTX_VIDEO_T2V_CONFIG,
}
def get_config_from_model_name(model_name: str) -> Dict[str, Any]:
if model_name not in SUPPORTED_MODEL_CONFIGS:
raise ValueError(
f"Model {model_name} not supported. Supported models are: {list(SUPPORTED_MODEL_CONFIGS.keys())}"
)
return SUPPORTED_MODEL_CONFIGS[model_name]
+23
View File
@@ -0,0 +1,23 @@
import torch
from accelerate import Accelerator
class State:
# Training state
seed: int = None
model_name: str = None
accelerator: Accelerator = None
weight_dtype: torch.dtype = None
train_epochs: int = None
train_steps: int = None
overwrote_max_train_steps: bool = False
num_trainable_parameters: int = 0
learning_rate: float = None
train_batch_size: int = None
generator: torch.Generator = None
# Hub state
repo_id: str = None
# Artifacts state
output_dir: str = None
+679
View File
@@ -0,0 +1,679 @@
import inspect
import json
import logging
import math
import os
import random
import shutil
from datetime import timedelta
from typing import Any, Dict
from pathlib import Path
import diffusers
import torch
import torch.backends
import transformers
import wandb
from accelerate import Accelerator, DistributedType
from accelerate.logging import get_logger
from accelerate.utils import (
DistributedDataParallelKwargs,
InitProcessGroupKwargs,
ProjectConfiguration,
set_seed,
gather_object,
)
from diffusers.optimization import get_scheduler
from diffusers.training_utils import (
cast_training_params,
compute_density_for_timestep_sampling,
compute_loss_weighting_for_sd3,
)
from diffusers.utils import export_to_video, load_image, load_video
from huggingface_hub import create_repo, upload_folder
from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from tqdm import tqdm
from .args import Args, validate_args
from .constants import FINETRAINERS_LOG_LEVEL
from .dataset import BucketSampler, VideoDatasetWithResizing
from .models import get_config_from_model_name
from .state import State
from .utils.file_utils import find_files, delete_files, string_to_filename
from .utils.optimizer_utils import get_optimizer, gradient_norm
from .utils.memory_utils import get_memory_statistics, free_memory, make_contiguous
from .utils.torch_utils import unwrap_model
logger = get_logger("finetrainers")
logger.setLevel(FINETRAINERS_LOG_LEVEL)
class Trainer:
def __init__(self, args: Args) -> None:
validate_args(args)
self.args = args
self.state = State()
# Tokenizers
self.tokenizer = None
self.tokenizer_2 = None
self.tokenizer_3 = None
# Text encoders
self.text_encoder = None
self.text_encoder_2 = None
self.text_encoder_3 = None
# Denoisers
self.transformer = None
self.unet = None
# Autoencoders
self.vae = None
self._init_distributed()
self._init_logging()
self._init_directories_and_repositories()
self.state.model_name = self.args.model_name
self.model_config = get_config_from_model_name(self.args.model_name)
def prepare_models(self) -> None:
logger.info("Initializing models")
# TODO(aryan): refactor in future
load_components_kwargs = {
"text_encoder_dtype": torch.bfloat16,
"transformer_dtype": torch.bfloat16,
"vae_dtype": torch.bfloat16,
"cache_dir": self.args.cache_dir,
}
if self.args.pretrained_model_name_or_path is not None:
load_components_kwargs["model_id"] = self.args.pretrained_model_name_or_path
components = self._model_config_call(self.model_config["load_components"], load_components_kwargs)
self.tokenizer = components.get("tokenizer", None)
self.text_encoder = components.get("text_encoder", None)
self.transformer = components.get("transformer", None)
self.vae = components.get("vae", None)
self.scheduler = components.get("scheduler", None)
self.transformer_config = self.transformer.config if self.transformer is not None else None
def prepare_dataset(self) -> None:
logger.info("Initializing dataset and dataloader")
self.dataset = VideoDatasetWithResizing(
data_root=self.args.data_root,
caption_column=self.args.caption_column,
video_column=self.args.video_column,
resolution_buckets=self.args.video_resolution_buckets,
dataset_file=self.args.dataset_file,
id_token=self.args.id_token,
)
self.dataloader = torch.utils.data.DataLoader(
self.dataset,
batch_size=1,
sampler=BucketSampler(self.dataset, batch_size=self.args.batch_size, shuffle=True),
collate_fn=self.model_config.get("collate_fn"),
num_workers=self.args.dataloader_num_workers,
pin_memory=self.args.pin_memory,
)
def prepare_trainable_parameters(self) -> None:
logger.info("Initializing trainable parameters")
# TODO(aryan): refactor later. for now only lora is supported
self.text_encoder.requires_grad_(False)
self.transformer.requires_grad_(False)
self.vae.requires_grad_(False)
# For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
weight_dtype = torch.float32
if self.state.accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
elif self.state.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."
)
# TODO(aryan): handle torch dtype from accelerator vs model dtype
self.state.weight_dtype = weight_dtype
self.text_encoder.to(self.state.accelerator.device, dtype=weight_dtype)
self.transformer.to(self.state.accelerator.device, dtype=weight_dtype)
self.vae.to(self.state.accelerator.device, dtype=weight_dtype)
if self.args.gradient_checkpointing:
self.transformer.enable_gradient_checkpointing()
transformer_lora_config = LoraConfig(
r=self.args.rank,
lora_alpha=self.args.lora_alpha,
init_lora_weights=True,
target_modules=self.args.target_modules,
)
self.transformer.add_adapter(transformer_lora_config)
# TODO: refactor
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
if self.state.accelerator.is_main_process:
transformer_lora_layers_to_save = None
for model in models:
if isinstance(
unwrap_model(self.state.accelerator, model),
type(unwrap_model(self.state.accelerator, self.transformer)),
):
model = unwrap_model(self.state.accelerator, 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()
self.model_config["pipeline_cls"].save_lora_weights(
output_dir,
transformer_lora_layers=transformer_lora_layers_to_save,
)
def load_model_hook(models, input_dir):
transformer_ = self.model_config["pipeline_cls"].from_pretrained(
self.args.pretrained_model_name_or_path, subfolder="transformer"
)
transformer_.add_adapter(transformer_lora_config)
lora_state_dict = self.model_config["pipeline_cls"].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.")
}
incompatible_keys = set_peft_model_state_dict(transformer_, transformer_state_dict, adapter_name="default")
if incompatible_keys is not None:
# check only for unexpected keys
unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None)
if unexpected_keys:
logger.warning(
f"Loading adapter weights from state_dict led to unexpected keys not found in the model: "
f" {unexpected_keys}. "
)
# Make sure the trainable params are in float32. This is again needed since the base models
# are in `weight_dtype`. More details:
# https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804
if self.args.mixed_precision == "fp16":
# only upcast trainable parameters (LoRA) into fp32
cast_training_params([transformer_])
self.state.accelerator.register_save_state_pre_hook(save_model_hook)
self.state.accelerator.register_load_state_pre_hook(load_model_hook)
# Enable TF32 for faster training on Ampere GPUs: https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
if self.args.allow_tf32 and torch.cuda.is_available():
torch.backends.cuda.matmul.allow_tf32 = True
def prepare_optimizer(self) -> None:
logger.info("Initializing optimizer and lr scheduler")
self.state.train_epochs = self.args.train_epochs
self.state.train_steps = self.args.train_steps
# Make sure the trainable params are in float32
if self.args.mixed_precision == "fp16":
# only upcast trainable parameters (LoRA) into fp32
cast_training_params([self.transformer], dtype=torch.float32)
self.state.learning_rate = self.args.lr
if self.args.scale_lr:
self.state.learning_rate = (
self.state.learning_rate
* self.args.gradient_accumulation_steps
* self.args.batch_size
* self.state.accelerator.num_processes
)
transformer_lora_parameters = list(filter(lambda p: p.requires_grad, self.transformer.parameters()))
transformer_parameters_with_lr = {
"params": transformer_lora_parameters,
"lr": self.state.learning_rate,
}
params_to_optimize = [transformer_parameters_with_lr]
self.state.num_trainable_parameters = sum(p.numel() for p in transformer_lora_parameters)
# TODO(aryan): add deepspeed support
optimizer = get_optimizer(
params_to_optimize=params_to_optimize,
optimizer_name=self.args.optimizer,
learning_rate=self.args.lr,
beta1=self.args.beta1,
beta2=self.args.beta2,
beta3=self.args.beta3,
epsilon=self.args.epsilon,
weight_decay=self.args.weight_decay,
)
num_update_steps_per_epoch = math.ceil(len(self.dataloader) / self.args.gradient_accumulation_steps)
if self.state.train_steps is None:
self.state.train_steps = self.state.train_epochs * num_update_steps_per_epoch
self.state.overwrote_max_train_steps = True
lr_scheduler = get_scheduler(
name=self.args.lr_scheduler,
optimizer=optimizer,
num_warmup_steps=self.args.lr_warmup_steps * self.state.accelerator.num_processes,
num_training_steps=self.state.train_steps * self.state.accelerator.num_processes,
num_cycles=self.args.lr_num_cycles,
power=self.args.lr_power,
)
self.optimizer = optimizer
self.lr_scheduler = lr_scheduler
def prepare_for_training(self) -> None:
self.transformer, self.optimizer, self.dataloader, self.lr_scheduler = self.state.accelerator.prepare(
self.transformer, self.optimizer, self.dataloader, self.lr_scheduler
)
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
num_update_steps_per_epoch = math.ceil(len(self.dataloader) / self.args.gradient_accumulation_steps)
if self.state.overwrote_max_train_steps:
self.state.train_steps = self.state.train_epochs * num_update_steps_per_epoch
# Afterwards we recalculate our number of training epochs
self.state.train_epochs = math.ceil(self.state.train_steps / num_update_steps_per_epoch)
def prepare_trackers(self) -> None:
logger.info("Initializing trackers")
tracker_name = self.args.tracker_name or "finetrainers-experiment"
self.state.accelerator.init_trackers(tracker_name, config=self.args.to_dict())
def train(self) -> None:
logger.info("Starting training")
memory_statistics = get_memory_statistics()
logger.info(f"Memory before training start: {json.dumps(memory_statistics, indent=4)}")
self.state.train_batch_size = (
self.args.batch_size * self.state.accelerator.num_processes * self.args.gradient_accumulation_steps
)
info = {
"trainable parameters": self.state.num_trainable_parameters,
"total samples": len(self.dataset),
"train epochs": self.state.train_epochs,
"train steps": self.state.train_steps,
"batches per device": self.args.batch_size,
"total batches observed per epoch": len(self.dataloader),
"train batch size": self.state.train_batch_size,
"gradient accumulation steps": self.args.gradient_accumulation_steps,
}
logger.info(f"Training configuration: {json.dumps(info, indent=4)}")
# TODO(aryan): handle resume from checkpoint
global_step = 0
first_epoch = 0
initial_global_step = 0
progress_bar = tqdm(
range(0, self.state.train_steps),
initial=initial_global_step,
desc="Training steps",
disable=not self.state.accelerator.is_local_main_process,
)
accelerator = self.state.accelerator
weight_dtype = self.state.weight_dtype
scheduler_sigmas = self.scheduler.sigmas.clone().to(device=accelerator.device, dtype=weight_dtype)
generator = torch.Generator(device=accelerator.device)
if self.args.seed is not None:
generator = generator.manual_seed(self.args.seed)
self.state.generator = generator
for epoch in range(first_epoch, self.state.train_epochs):
logger.debug(f"Starting epoch ({epoch + 1}/{self.state.train_epochs})")
self.transformer.train()
models_to_accumulate = [self.transformer]
for step, batch in enumerate(self.dataloader):
logger.debug(f"Starting step {step + 1}")
logs = {}
with accelerator.accumulate(models_to_accumulate):
videos = batch["videos"]
prompts = batch["prompts"]
batch_size = len(prompts)
if self.args.caption_dropout_technique == "empty":
if random.random() < self.args.caption_dropout_p:
prompts = [""] * batch_size
latent_conditions = self.model_config["prepare_latents"](
vae=self.vae,
image_or_video=videos,
patch_size=self.transformer_config.patch_size,
patch_size_t=self.transformer_config.patch_size_t,
device=accelerator.device,
dtype=weight_dtype,
generator=generator,
)
latent_conditions = make_contiguous(latent_conditions)
other_conditions = self.model_config["prepare_conditions"](
tokenizer=self.tokenizer,
text_encoder=self.text_encoder,
prompt=prompts,
device=accelerator.device,
dtype=weight_dtype,
)
other_conditions = make_contiguous(other_conditions)
if self.args.caption_dropout_technique == "zero":
if random.random() < self.args.caption_dropout_p:
other_conditions["prompt_embeds"].fill_(0)
other_conditions["prompt_attention_mask"].fill_(False)
# These weighting schemes use a uniform timestep sampling and instead post-weight the loss
weights = compute_density_for_timestep_sampling(
weighting_scheme=self.args.flow_weighting_scheme,
batch_size=batch_size,
logit_mean=self.args.flow_logit_mean,
logit_std=self.args.flow_logit_std,
mode_scale=self.args.flow_mode_scale,
)
indices = (weights * self.scheduler.config.num_train_timesteps).long()
sigmas = scheduler_sigmas[indices].flatten()
while sigmas.ndim < latent_conditions["latents"].ndim:
sigmas = sigmas.unsqueeze(-1)
timesteps = (sigmas * 1000.0).long()
noise = torch.randn(
latent_conditions["latents"].shape,
generator=generator,
device=accelerator.device,
dtype=weight_dtype,
)
noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise
latent_conditions.update({"noisy_latents": noisy_latents})
other_conditions.update({"timesteps": timesteps})
# These weighting schemes use a uniform timestep sampling and instead post-weight the loss
weights = compute_loss_weighting_for_sd3(
weighting_scheme=self.args.flow_weighting_scheme, sigmas=sigmas
)
pred = self.model_config["forward_pass"](
transformer=self.transformer, **latent_conditions, **other_conditions
)
target = noise - latent_conditions["latents"]
loss = weights.float() * (pred["latents"].float() - target.float()).pow(2)
# Average loss across channel dimension
loss = loss.mean(list(range(1, loss.ndim)))
# Average loss across batch dimension
loss = loss.mean()
accelerator.backward(loss)
if accelerator.sync_gradients and accelerator.distributed_type != DistributedType.DEEPSPEED:
accelerator.clip_grad_norm_(self.transformer.parameters(), self.args.max_grad_norm)
self.optimizer.step()
self.lr_scheduler.step()
self.optimizer.zero_grad()
# Checks if the accelerator has performed an optimization step behind the scenes
if accelerator.sync_gradients:
progress_bar.update(1)
global_step += 1
# Checkpointing
if accelerator.distributed_type == DistributedType.DEEPSPEED or accelerator.is_main_process:
if global_step % self.args.checkpointing_steps == 0:
# before saving state, check if this save would set us over the `checkpointing_limit`
if self.args.checkpointing_limit is not None:
checkpoints = find_files(self.args.output_dir, prefix="checkpoint")
# before we save the new checkpoint, we need to have at_most `checkpoints_total_limit - 1` checkpoints
if len(checkpoints) >= self.args.checkpointing_limit:
num_to_remove = len(checkpoints) - self.args.checkpointing_limit + 1
checkpoints_to_remove = checkpoints[0:num_to_remove]
delete_files(checkpoints_to_remove)
logger.info(f"Checkpointing at step {global_step}")
save_path = os.path.join(self.args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(save_path)
logger.info(f"Saved state to {save_path}")
# Maybe run validation
should_run_validation = (
self.args.validation_every_n_steps is not None
and global_step % self.args.validation_every_n_steps == 0
)
if should_run_validation:
self.validate(global_step)
logs = {"loss": loss.detach().item(), "lr": self.lr_scheduler.get_last_lr()[0]}
progress_bar.set_postfix(logs)
accelerator.log(logs, step=global_step)
if global_step >= self.state.train_steps:
break
memory_statistics = get_memory_statistics()
logger.info(f"Memory after epoch {epoch + 1}: {json.dumps(memory_statistics, indent=4)}")
# Maybe run validation
should_run_validation = (
self.args.validation_every_n_epochs is not None
and (epoch + 1) % self.args.validation_every_n_epochs == 0
)
if should_run_validation:
self.validate(global_step)
accelerator.wait_for_everyone()
if accelerator.is_main_process:
self.transformer = unwrap_model(accelerator, self.transformer)
dtype = (
torch.float16
if self.args.mixed_precision == "fp16"
else torch.bfloat16
if self.args.mixed_precision == "bf16"
else torch.float32
)
self.transformer = self.transformer.to(dtype)
transformer_lora_layers = get_peft_model_state_dict(self.transformer)
self.model_config["pipeline_cls"].save_lora_weights(
save_directory=self.args.output_dir,
transformer_lora_layers=transformer_lora_layers,
)
del self.tokenizer, self.text_encoder, self.transformer, self.vae, self.scheduler
free_memory()
memory_statistics = get_memory_statistics()
logger.info(f"Memory after training end: {json.dumps(memory_statistics, indent=4)}")
accelerator.end_training()
def validate(self, step: int) -> None:
logger.info("Starting validation")
accelerator = self.state.accelerator
num_validation_samples = len(self.args.validation_prompts)
if num_validation_samples == 0:
logger.warning("No validation samples found. Skipping validation.")
return
self.transformer.eval()
memory_statistics = get_memory_statistics()
logger.info(f"Memory before validation start: {json.dumps(memory_statistics, indent=4)}")
pipeline = self.model_config["initialize_pipeline"](
model_id=self.args.pretrained_model_name_or_path,
cache_dir=self.args.cache_dir,
tokenizer=self.tokenizer,
text_encoder=self.text_encoder,
transformer=unwrap_model(accelerator, self.transformer),
vae=self.vae,
device=accelerator.device,
enable_slicing=self.args.enable_slicing,
enable_tiling=self.args.enable_tiling,
enable_model_cpu_offload=self.args.enable_model_cpu_offload,
)
all_processes_artifacts = []
for i in range(num_validation_samples):
# Skip current validation on all processes but one
if i % accelerator.num_processes != accelerator.process_index:
continue
prompt = self.args.validation_prompts[i]
image = self.args.validation_images[i]
video = self.args.validation_videos[i]
height = self.args.validation_heights[i]
width = self.args.validation_widths[i]
num_frames = self.args.validation_num_frames[i]
if image is not None:
image = load_image(image)
if video is not None:
video = load_video(video)
logger.debug(
f"Validating sample {i + 1}/{num_validation_samples} on process {accelerator.process_index}. Prompt: {prompt}",
main_process_only=False,
)
validation_artifacts = self.model_config["validation"](
pipeline=pipeline,
prompt=prompt,
image=image,
video=video,
height=height,
width=width,
num_frames=num_frames,
num_videos_per_prompt=self.args.num_validation_videos_per_prompt,
generator=self.state.generator,
)
prompt_filename = string_to_filename(prompt)[:25]
artifacts = {
"image": {"type": "image", "value": image},
"video": {"type": "video", "value": video},
}
for i, (artifact_type, artifact_value) in enumerate(validation_artifacts):
artifacts.update({f"artifact_{i}": {"type": artifact_type, "value": artifact_value}})
logger.debug(
f"Validation artifacts on process {accelerator.process_index}: {list(artifacts.keys())}",
main_process_only=False,
)
for key, value in list(artifacts.items()):
artifact_type = value["type"]
artifact_value = value["value"]
if artifact_type not in ["image", "video"] or artifact_value is None:
continue
extension = "png" if artifact_type == "image" else "mp4"
filename = f"validation-{step}-{accelerator.process_index}-{prompt_filename}.{extension}"
filename = os.path.join(self.args.output_dir, filename)
if artifact_type == "image":
logger.debug(f"Saving image to {filename}")
artifact_value.save(filename)
artifact_value = wandb.Image(filename)
elif artifact_type == "video":
logger.debug(f"Saving video to {filename}")
export_to_video(artifact_value, filename, fps=15)
artifact_value = wandb.Video(filename, caption=prompt)
all_processes_artifacts.append(artifact_value)
all_artifacts = gather_object(all_processes_artifacts)
if accelerator.is_main_process:
for tracker in accelerator.trackers:
if tracker.name == "wandb":
tracker.log({"validation": all_artifacts}, step=step)
accelerator.wait_for_everyone()
free_memory()
memory_statistics = get_memory_statistics()
logger.info(f"Memory after validation end: {json.dumps(memory_statistics, indent=4)}")
self.transformer.train()
def evaluate(self) -> None:
logger.info("Starting evaluation")
# TODO: implement metrics for evaluation
def _init_distributed(self) -> None:
logging_dir = Path(self.args.output_dir, self.args.logging_dir)
project_config = ProjectConfiguration(project_dir=self.args.output_dir, logging_dir=logging_dir)
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
init_process_group_kwargs = InitProcessGroupKwargs(
backend="nccl", timeout=timedelta(seconds=self.args.nccl_timeout)
)
mixed_precision = "no" if torch.backends.mps.is_available() else self.args.mixed_precision
report_to = None if self.args.report_to.lower() == "none" else self.args.report_to
accelerator = Accelerator(
project_config=project_config,
gradient_accumulation_steps=self.args.gradient_accumulation_steps,
mixed_precision=mixed_precision,
log_with=report_to,
kwargs_handlers=[ddp_kwargs, init_process_group_kwargs],
)
# Disable AMP for MPS.
if torch.backends.mps.is_available():
accelerator.native_amp = False
self.state.accelerator = accelerator
if self.args.seed is not None:
self.state.seed = self.args.seed
set_seed(self.args.seed)
def _init_logging(self) -> None:
logging.basicConfig(
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
datefmt="%m/%d/%Y %H:%M:%S",
level=FINETRAINERS_LOG_LEVEL,
)
if self.state.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()
logger.info("Initialized FineTrainers")
logger.info(self.state.accelerator.state, main_process_only=False)
def _init_directories_and_repositories(self) -> None:
if self.state.accelerator.is_main_process:
self.args.output_dir = Path(self.args.output_dir)
self.args.output_dir.mkdir(parents=True, exist_ok=True)
self.state.output_dir = self.args.output_dir
if self.args.push_to_hub:
repo_id = self.args.hub_model_id or Path(self.args.output_dir).name
self.state.repo_id = create_repo(token=self.args.hub_token, name=repo_id).repo_id
def _model_config_call(self, fn, kwargs):
accepted_kwargs = inspect.signature(fn).parameters.keys()
kwargs = {k: v for k, v in kwargs.items() if k in accepted_kwargs}
return fn(**kwargs)
+5
View File
@@ -0,0 +1,5 @@
from .file_utils import find_files, delete_files
from .diffusion_utils import resolution_dependant_timestep_flow_shift, default_flow_shift
from .memory_utils import get_memory_statistics, bytes_to_gigabytes, free_memory, make_contiguous
from .torch_utils import unwrap_model
from .optimizer_utils import get_optimizer, gradient_norm, max_gradient
+30
View File
@@ -0,0 +1,30 @@
import torch
# Default values copied from https://github.com/huggingface/diffusers/blob/8957324363d8b239d82db4909fbf8c0875683e3d/src/diffusers/schedulers/scheduling_flow_match_euler_discrete.py#L47
def resolution_dependant_timestep_flow_shift(
latents: torch.Tensor,
sigmas: torch.Tensor,
base_image_seq_len: int = 256,
max_image_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15,
) -> torch.Tensor:
image_or_video_sequence_length = 0
if latents.ndim == 4:
image_or_video_sequence_length = latents.shape[2] * latents.shape[3]
elif latents.ndim == 5:
image_or_video_sequence_length = latents.shape[2] * latents.shape[3] * latents.shape[4]
else:
raise ValueError(f"Expected 4D or 5D tensor, got {latents.ndim}D tensor")
m = (max_shift - base_shift) / (max_image_seq_len - base_image_seq_len)
b = base_shift - m * base_image_seq_len
mu = m * image_or_video_sequence_length + b
sigmas = default_flow_shift(latents, sigmas, shift=mu)
return sigmas
def default_flow_shift(sigmas: torch.Tensor, shift: float = 1.0) -> torch.Tensor:
sigmas = (sigmas * shift) / (1 + (shift - 1) * sigmas)
return sigmas
+44
View File
@@ -0,0 +1,44 @@
import logging
import os
import shutil
from pathlib import Path
from typing import Any, Dict, List, Union
logger = logging.getLogger("finetrainers")
logger.setLevel(os.environ.get("FINETRAINERS_LOG_LEVEL", "INFO"))
def find_files(dir: Union[str, Path], prefix: str = "checkpoint") -> List[str]:
if not isinstance(dir, Path):
dir = Path(dir)
if not dir.exists():
return []
checkpoints = os.listdir(dir.as_posix())
checkpoints = [c for c in checkpoints if c.startswith(prefix)]
checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1]))
return checkpoints
def delete_files(dirs: Union[str, List[str], Path, List[Path]]) -> None:
if not isinstance(dirs, list):
dirs = [dirs]
dirs = [Path(d) if isinstance(d, str) else d for d in dirs]
logger.info(f"Deleting files: {dirs}")
for dir in dirs:
if not dir.exists():
continue
shutil.rmtree(dir, ignore_errors=True)
def string_to_filename(s: str) -> str:
return (
s.replace(" ", "-")
.replace("/", "-")
.replace(":", "-")
.replace(".", "-")
.replace(",", "-")
.replace(";", "-")
.replace("!", "-")
.replace("?", "-")
)
+59
View File
@@ -0,0 +1,59 @@
import gc
import logging
from typing import Any, Dict, Union
import torch
from accelerate.logging import get_logger
logger = get_logger("finetrainers")
def get_memory_statistics(precision: int = 3) -> Dict[str, Any]:
memory_allocated = None
memory_reserved = None
max_memory_allocated = None
max_memory_reserved = None
if torch.cuda.is_available():
device = torch.cuda.current_device()
memory_allocated = torch.cuda.memory_allocated(device)
memory_reserved = torch.cuda.memory_reserved(device)
max_memory_allocated = torch.cuda.max_memory_allocated(device)
max_memory_reserved = torch.cuda.max_memory_reserved(device)
elif torch.mps.is_available():
memory_allocated = torch.mps.current_allocated_memory()
else:
logger.warning("No CUDA, MPS, or ROCm device found. Memory statistics are not available.")
return {
"memory_allocated": round(bytes_to_gigabytes(memory_allocated), ndigits=precision),
"memory_reserved": round(bytes_to_gigabytes(memory_reserved), ndigits=precision),
"max_memory_allocated": round(bytes_to_gigabytes(max_memory_allocated), ndigits=precision),
"max_memory_reserved": round(bytes_to_gigabytes(max_memory_reserved), ndigits=precision),
}
def bytes_to_gigabytes(x: int) -> float:
if x is not None:
return x / 1024**3
def free_memory() -> None:
if torch.cuda.is_available():
gc.collect()
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
# TODO(aryan): handle non-cuda devices
def make_contiguous(x: Union[torch.Tensor, Dict[str, torch.Tensor]]) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if isinstance(x, torch.Tensor):
return x.contiguous()
elif isinstance(x, dict):
return {k: make_contiguous(v) for k, v in x.items()}
else:
return x
+178
View File
@@ -0,0 +1,178 @@
import inspect
import logging
from accelerate.logging import get_logger
import torch
logger = get_logger("finetrainers")
def get_optimizer(
params_to_optimize,
optimizer_name: str = "adam",
learning_rate: float = 1e-3,
beta1: float = 0.9,
beta2: float = 0.95,
beta3: float = 0.98,
epsilon: float = 1e-8,
weight_decay: float = 1e-4,
prodigy_decouple: bool = False,
prodigy_use_bias_correction: bool = False,
prodigy_safeguard_warmup: bool = False,
use_8bit: bool = False,
use_4bit: bool = False,
use_torchao: bool = False,
use_deepspeed: bool = False,
use_cpu_offload_optimizer: bool = False,
offload_gradients: bool = False,
) -> torch.optim.Optimizer:
optimizer_name = optimizer_name.lower()
# Use DeepSpeed optimzer
if use_deepspeed:
from accelerate.utils import DummyOptim
return DummyOptim(
params_to_optimize,
lr=learning_rate,
betas=(beta1, beta2),
eps=epsilon,
weight_decay=weight_decay,
)
if use_8bit and use_4bit:
raise ValueError("Cannot set both `use_8bit` and `use_4bit` to True.")
if (use_torchao and (use_8bit or use_4bit)) or use_cpu_offload_optimizer:
try:
import torchao
torchao.__version__
except ImportError:
raise ImportError(
"To use optimizers from torchao, please install the torchao library: `USE_CPP=0 pip install torchao`."
)
if not use_torchao and use_4bit:
raise ValueError("4-bit Optimizers are only supported with torchao.")
# Optimizer creation
supported_optimizers = ["adam", "adamw", "prodigy", "came"]
if optimizer_name not in supported_optimizers:
logger.warning(
f"Unsupported choice of optimizer: {optimizer_name}. Supported optimizers include {supported_optimizers}. Defaulting to `AdamW`."
)
optimizer_name = "adamw"
if (use_8bit or use_4bit) and optimizer_name not in ["adam", "adamw"]:
raise ValueError("`use_8bit` and `use_4bit` can only be used with the Adam and AdamW optimizers.")
if use_8bit:
try:
import bitsandbytes as bnb
except ImportError:
raise ImportError(
"To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`."
)
if optimizer_name == "adamw":
if use_torchao:
from torchao.prototype.low_bit_optim import AdamW4bit, AdamW8bit
optimizer_class = AdamW8bit if use_8bit else AdamW4bit if use_4bit else torch.optim.AdamW
else:
optimizer_class = bnb.optim.AdamW8bit if use_8bit else torch.optim.AdamW
init_kwargs = {
"betas": (beta1, beta2),
"eps": epsilon,
"weight_decay": weight_decay,
}
elif optimizer_name == "adam":
if use_torchao:
from torchao.prototype.low_bit_optim import Adam4bit, Adam8bit
optimizer_class = Adam8bit if use_8bit else Adam4bit if use_4bit else torch.optim.Adam
else:
optimizer_class = bnb.optim.Adam8bit if use_8bit else torch.optim.Adam
init_kwargs = {
"betas": (beta1, beta2),
"eps": epsilon,
"weight_decay": weight_decay,
}
elif optimizer_name == "prodigy":
try:
import prodigyopt
except ImportError:
raise ImportError("To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`")
optimizer_class = prodigyopt.Prodigy
if learning_rate <= 0.1:
logger.warning(
"Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0"
)
init_kwargs = {
"lr": learning_rate,
"betas": (beta1, beta2),
"beta3": beta3,
"eps": epsilon,
"weight_decay": weight_decay,
"decouple": prodigy_decouple,
"use_bias_correction": prodigy_use_bias_correction,
"safeguard_warmup": prodigy_safeguard_warmup,
}
elif optimizer_name == "came":
try:
import came_pytorch
except ImportError:
raise ImportError("To use CAME, please install the came-pytorch library: `pip install came-pytorch`")
optimizer_class = came_pytorch.CAME
init_kwargs = {
"lr": learning_rate,
"eps": (1e-30, 1e-16),
"betas": (beta1, beta2, beta3),
"weight_decay": weight_decay,
}
if use_cpu_offload_optimizer:
from torchao.prototype.low_bit_optim import CPUOffloadOptimizer
if "fused" in inspect.signature(optimizer_class.__init__).parameters:
init_kwargs.update({"fused": True})
optimizer = CPUOffloadOptimizer(
params_to_optimize, optimizer_class=optimizer_class, offload_gradients=offload_gradients, **init_kwargs
)
else:
optimizer = optimizer_class(params_to_optimize, **init_kwargs)
return optimizer
def gradient_norm(parameters):
norm = 0
for param in parameters:
if param.grad is None:
continue
local_norm = param.grad.detach().data.norm(2)
norm += local_norm.item() ** 2
norm = norm**0.5
return norm
def max_gradient(parameters):
max_grad_value = float("-inf")
for param in parameters:
if param.grad is None:
continue
local_max_grad = param.grad.detach().data.abs().max()
max_grad_value = max(max_grad_value, local_max_grad.item())
return max_grad_value
+8
View File
@@ -0,0 +1,8 @@
from accelerate import Accelerator
from diffusers.utils.torch_utils import is_compiled_module
def unwrap_model(accelerator: Accelerator, model):
model = accelerator.unwrap_model(model)
model = model._orig_mod if is_compiled_module(model) else model
return model
+45
View File
@@ -0,0 +1,45 @@
import logging
import traceback
from finetrainers import Trainer, parse_arguments
from finetrainers.constants import FINETRAINERS_LOG_LEVEL
logger = logging.getLogger("finetrainers")
logger.setLevel(FINETRAINERS_LOG_LEVEL)
def main():
try:
import multiprocessing
multiprocessing.set_start_method("fork")
except Exception as e:
logger.error(
f'Failed to set multiprocessing start method to "fork". This can lead to poor performance, high memory usage, or crashes. '
f"See: https://pytorch.org/docs/stable/notes/multiprocessing.html\n"
f"Error: {e}"
)
try:
args = parse_arguments()
trainer = Trainer(args)
trainer.prepare_dataset()
trainer.prepare_models()
trainer.prepare_trainable_parameters()
trainer.prepare_optimizer()
trainer.prepare_for_training()
trainer.prepare_trackers()
trainer.train()
trainer.evaluate()
except KeyboardInterrupt:
logger.info("Received keyboard interrupt. Exiting...")
except Exception as e:
logger.error(f"An error occurred during training: {e}")
logger.error(traceback.format_exc())
if __name__ == "__main__":
main()
+1 -1
View File
@@ -30,7 +30,7 @@ for learning_rate in "${LEARNING_RATES[@]}"; do
for steps in "${MAX_TRAIN_STEPS[@]}"; do
output_dir="/path/to/my/models/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox_image_to_video_lora.py \
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox/cogvideox_image_to_video_lora.py \
--pretrained_model_name_or_path THUDM/CogVideoX-5b-I2V \
--data_root $DATA_ROOT \
--caption_column $CAPTION_COLUMN \
+1 -1
View File
@@ -37,7 +37,7 @@ for learning_rate in "${LEARNING_RATES[@]}"; do
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE \
--gpu_ids $GPU_IDS \
training/cogvideox_image_to_video_sft.py \
training/cogvideox/cogvideox_image_to_video_sft.py \
--pretrained_model_name_or_path $MODEL_PATH \
--data_root $DATA_ROOT \
--caption_column $CAPTION_COLUMN \
+1 -1
View File
@@ -34,7 +34,7 @@ for learning_rate in "${LEARNING_RATES[@]}"; do
for steps in "${MAX_TRAIN_STEPS[@]}"; do
output_dir="./cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox_text_to_video_lora.py \
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox/cogvideox_text_to_video_lora.py \
--pretrained_model_name_or_path $MODEL_PATH \
--data_root $DATA_ROOT \
--caption_column $CAPTION_COLUMN \
+1 -1
View File
@@ -30,7 +30,7 @@ for learning_rate in "${LEARNING_RATES[@]}"; do
for steps in "${MAX_TRAIN_STEPS[@]}"; do
output_dir="/path/to/my/models/cogvideox-sft__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox_text_to_video_sft.py \
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox/cogvideox_text_to_video_sft.py \
--pretrained_model_name_or_path THUDM/CogVideoX-5b \
--data_root $DATA_ROOT \
--caption_column $CAPTION_COLUMN \
+459
View File
@@ -0,0 +1,459 @@
# CogVideoX Factory 🧪
[中文阅读](./README_zh.md)
Fine-tune Cog family of video models for custom video generation under 24GB of GPU memory ⚡️📼
<table align="center">
<tr>
<td align="center"><video src="https://github.com/user-attachments/assets/aad07161-87cb-4784-9e6b-16d06581e3e5">Your browser does not support the video tag.</video></td>
</tr>
</table>
**Update 29 Nov 2024**: We have added an experimental memory-efficient trainer for Mochi-1. Check it out [here](https://github.com/a-r-r-o-w/cogvideox-factory/blob/main/training/mochi-1/)!
## Quickstart
Clone the repository and make sure the requirements are installed: `pip install -r requirements.txt` and install diffusers from source by `pip install git+https://github.com/huggingface/diffusers`.
Then download a dataset:
```bash
# install `huggingface_hub`
huggingface-cli download \
--repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset \
--local-dir video-dataset-disney
```
Then launch LoRA fine-tuning for text-to-video (modify the different hyperparameters, dataset root, and other configuration options as per your choice):
```bash
# For LoRA finetuning of the text-to-video CogVideoX models
./train_text_to_video_lora.sh
# For full finetuning of the text-to-video CogVideoX models
./train_text_to_video_sft.sh
# For LoRA finetuning of the image-to-video CogVideoX models
./train_image_to_video_lora.sh
```
Assuming your LoRA is saved and pushed to the HF Hub, and named `my-awesome-name/my-awesome-lora`, we can now use the finetuned model for inference:
```diff
import torch
from diffusers import CogVideoXPipeline
from diffusers.utils import export_to_video
pipe = CogVideoXPipeline.from_pretrained(
"THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16
).to("cuda")
+ pipe.load_lora_weights("my-awesome-name/my-awesome-lora", adapter_name="cogvideox-lora")
+ pipe.set_adapters(["cogvideox-lora"], [1.0])
video = pipe("<my-awesome-prompt>").frames[0]
export_to_video(video, "output.mp4", fps=8)
```
For Image-to-Video LoRAs trained with multiresolution videos, one must also add the following lines (see [this](https://github.com/a-r-r-o-w/cogvideox-factory/issues/26) Issue for more details):
```python
from diffusers import CogVideoXImageToVideoPipeline
pipe = CogVideoXImageToVideoPipeline.from_pretrained(
"THUDM/CogVideoX-5b-I2V", torch_dtype=torch.bfloat16
).to("cuda")
# ...
del pipe.transformer.patch_embed.pos_embedding
pipe.transformer.patch_embed.use_learned_positional_embeddings = False
pipe.transformer.config.use_learned_positional_embeddings = False
```
You can also check if your LoRA is correctly mounted [here](tests/test_lora_inference.py).
Below we provide additional sections detailing on more options explored in this repository. They all attempt to make fine-tuning for video models as accessible as possible by reducing memory requirements as much as possible.
## Prepare Dataset and Training
Before starting the training, please check whether the dataset has been prepared according to the [dataset specifications](assets/dataset.md). We provide training scripts suitable for text-to-video and image-to-video generation, compatible with the [CogVideoX model family](https://huggingface.co/collections/THUDM/cogvideo-66c08e62f1685a3ade464cce). Training can be started using the `train*.sh` scripts, depending on the task you want to train. Let's take LoRA fine-tuning for text-to-video as an example.
- Configure environment variables as per your choice:
```bash
export TORCH_LOGS="+dynamo,recompiles,graph_breaks"
export TORCHDYNAMO_VERBOSE=1
export WANDB_MODE="offline"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
```
- Configure which GPUs to use for training: `GPU_IDS="0,1"`
- Choose hyperparameters for training. Let's try to do a sweep on learning rate and optimizer type as an example:
```bash
LEARNING_RATES=("1e-4" "1e-3")
LR_SCHEDULES=("cosine_with_restarts")
OPTIMIZERS=("adamw" "adam")
MAX_TRAIN_STEPS=("3000")
```
- Select which Accelerate configuration you would like to train with: `ACCELERATE_CONFIG_FILE="accelerate_configs/uncompiled_1.yaml"`. We provide some default configurations in the `accelerate_configs/` directory - single GPU uncompiled/compiled, 2x GPU DDP, DeepSpeed, etc. You can create your own config files with custom settings using `accelerate config --config_file my_config.yaml`.
- Specify the absolute paths and columns/files for captions and videos.
```bash
DATA_ROOT="/path/to/my/datasets/video-dataset-disney"
CAPTION_COLUMN="prompt.txt"
VIDEO_COLUMN="videos.txt"
```
- Launch experiments sweeping different hyperparameters:
```
for learning_rate in "${LEARNING_RATES[@]}"; do
for lr_schedule in "${LR_SCHEDULES[@]}"; do
for optimizer in "${OPTIMIZERS[@]}"; do
for steps in "${MAX_TRAIN_STEPS[@]}"; do
output_dir="/path/to/my/models/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox/cogvideox_text_to_video_lora.py \
--pretrained_model_name_or_path THUDM/CogVideoX-5b \
--data_root $DATA_ROOT \
--caption_column $CAPTION_COLUMN \
--video_column $VIDEO_COLUMN \
--id_token BW_STYLE \
--height_buckets 480 \
--width_buckets 720 \
--frame_buckets 49 \
--dataloader_num_workers 8 \
--pin_memory \
--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:::BW_STYLE A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance\" \
--validation_prompt_separator ::: \
--num_validation_videos 1 \
--validation_epochs 10 \
--seed 42 \
--rank 128 \
--lora_alpha 128 \
--mixed_precision bf16 \
--output_dir $output_dir \
--max_num_frames 49 \
--train_batch_size 1 \
--max_train_steps $steps \
--checkpointing_steps 1000 \
--gradient_accumulation_steps 1 \
--gradient_checkpointing \
--learning_rate $learning_rate \
--lr_scheduler $lr_schedule \
--lr_warmup_steps 400 \
--lr_num_cycles 1 \
--enable_slicing \
--enable_tiling \
--optimizer $optimizer \
--beta1 0.9 \
--beta2 0.95 \
--weight_decay 0.001 \
--max_grad_norm 1.0 \
--allow_tf32 \
--report_to wandb \
--nccl_timeout 1800"
echo "Running command: $cmd"
eval $cmd
echo -ne "-------------------- Finished executing script --------------------\n\n"
done
done
done
done
```
To understand what the different parameters mean, you could either take a look at the [args](./training/args.py) file or run the training script with `--help`.
Note: Training scripts are untested on MPS, so performance and memory requirements can differ widely compared to the CUDA reports below.
## Memory requirements
<table align="center">
<tr>
<td align="center" colspan="2"><b>CogVideoX LoRA Finetuning</b></td>
</tr>
<tr>
<td align="center"><a href="https://huggingface.co/THUDM/CogVideoX-2b">THUDM/CogVideoX-2b</a></td>
<td align="center"><a href="https://huggingface.co/THUDM/CogVideoX-5b">THUDM/CogVideoX-5b</a></td>
</tr>
<tr>
<td align="center"><img src="assets/lora_2b.png" /></td>
<td align="center"><img src="assets/lora_5b.png" /></td>
</tr>
<tr>
<td align="center" colspan="2"><b>CogVideoX Full Finetuning</b></td>
</tr>
<tr>
<td align="center"><a href="https://huggingface.co/THUDM/CogVideoX-2b">THUDM/CogVideoX-2b</a></td>
<td align="center"><a href="https://huggingface.co/THUDM/CogVideoX-5b">THUDM/CogVideoX-5b</a></td>
</tr>
<tr>
<td align="center"><img src="assets/sft_2b.png" /></td>
<td align="center"><img src="assets/sft_5b.png" /></td>
</tr>
</table>
Supported and verified memory optimizations for training include:
- `CPUOffloadOptimizer` from [`torchao`](https://github.com/pytorch/ao). You can read about its capabilities and limitations [here](https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim#optimizer-cpu-offload). In short, it allows you to use the CPU for storing trainable parameters and gradients. This results in the optimizer step happening on the CPU, which requires a fast CPU optimizer, such as `torch.optim.AdamW(fused=True)` or applying `torch.compile` on the optimizer step. Additionally, it is recommended not to `torch.compile` your model for training. Gradient clipping and accumulation is not supported yet either.
- Low-bit optimizers from [`bitsandbytes`](https://huggingface.co/docs/bitsandbytes/optimizers). TODO: to test and make [`torchao`](https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim) ones work
- DeepSpeed Zero2: Since we rely on `accelerate`, follow [this guide](https://huggingface.co/docs/accelerate/en/usage_guides/deepspeed) to configure your `accelerate` installation to enable training with DeepSpeed Zero2 optimizations.
> [!IMPORTANT]
> The memory requirements are reported after running the `training/prepare_dataset.py`, which converts the videos and captions to latents and embeddings. During training, we directly load the latents and embeddings, and do not require the VAE or the T5 text encoder. However, if you perform validation/testing, these must be loaded and increase the amount of required memory. Not performing validation/testing saves a significant amount of memory, which can be used to focus solely on training if you're on smaller VRAM GPUs.
>
> If you choose to run validation/testing, you can save some memory on lower VRAM GPUs by specifying `--enable_model_cpu_offload`.
### LoRA finetuning
> [!NOTE]
> The memory requirements for image-to-video lora finetuning are similar to that of text-to-video on `THUDM/CogVideoX-5b`, so it hasn't been reported explicitly.
>
> Additionally, to prepare test images for I2V finetuning, you could either generate them on-the-fly by modifying the script, or extract some frames from your training data using:
> `ffmpeg -i input.mp4 -frames:v 1 frame.png`,
> or provide a URL to a valid and accessible image.
<details>
<summary> AdamW </summary>
**Note:** Trying to run CogVideoX-5b without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
With `train_batch_size = 1`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 16 | False | 12.945 | 43.764 | 46.918 | 24.234 |
| THUDM/CogVideoX-2b | 16 | True | 12.945 | 12.945 | 21.121 | 24.234 |
| THUDM/CogVideoX-2b | 64 | False | 13.035 | 44.314 | 47.469 | 24.469 |
| THUDM/CogVideoX-2b | 64 | True | 13.036 | 13.035 | 21.564 | 24.500 |
| THUDM/CogVideoX-2b | 256 | False | 13.095 | 45.826 | 48.990 | 25.543 |
| THUDM/CogVideoX-2b | 256 | True | 13.094 | 13.095 | 22.344 | 25.537 |
| THUDM/CogVideoX-5b | 16 | True | 19.742 | 19.742 | 28.746 | 38.123 |
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 20.818 | 30.338 | 38.738 |
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 22.119 | 31.939 | 41.537 |
With `train_batch_size = 4`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 16 | True | 12.945 | 21.803 | 21.814 | 24.322 |
| THUDM/CogVideoX-2b | 64 | True | 13.035 | 22.254 | 22.254 | 24.572 |
| THUDM/CogVideoX-2b | 256 | True | 13.094 | 22.020 | 22.033 | 25.574 |
| THUDM/CogVideoX-5b | 16 | True | 19.742 | 46.492 | 46.492 | 38.197 |
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 47.805 | 47.805 | 39.365 |
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 47.268 | 47.332 | 41.008 |
</details>
<details>
<summary> AdamW (8-bit bitsandbytes) </summary>
**Note:** Trying to run CogVideoX-5b without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
With `train_batch_size = 1`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 16 | False | 12.945 | 43.732 | 46.887 | 24.195 |
| THUDM/CogVideoX-2b | 16 | True | 12.945 | 12.945 | 21.430 | 24.195 |
| THUDM/CogVideoX-2b | 64 | False | 13.035 | 44.004 | 47.158 | 24.369 |
| THUDM/CogVideoX-2b | 64 | True | 13.035 | 13.035 | 21.297 | 24.357 |
| THUDM/CogVideoX-2b | 256 | False | 13.035 | 45.291 | 48.455 | 24.836 |
| THUDM/CogVideoX-2b | 256 | True | 13.035 | 13.035 | 21.625 | 24.869 |
| THUDM/CogVideoX-5b | 16 | True | 19.742 | 19.742 | 28.602 | 38.049 |
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 20.818 | 29.359 | 38.520 |
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 21.352 | 30.727 | 39.596 |
With `train_batch_size = 4`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 16 | True | 12.945 | 21.734 | 21.775 | 24.281 |
| THUDM/CogVideoX-2b | 64 | True | 13.036 | 21.941 | 21.941 | 24.445 |
| THUDM/CogVideoX-2b | 256 | True | 13.094 | 22.020 | 22.266 | 24.943 |
| THUDM/CogVideoX-5b | 16 | True | 19.742 | 46.320 | 46.326 | 38.104 |
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 46.820 | 46.820 | 38.588 |
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 47.920 | 47.980 | 40.002 |
</details>
<details>
<summary> AdamW + CPUOffloadOptimizer (with gradient offloading) </summary>
**Note:** Trying to run CogVideoX-5b without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
With `train_batch_size = 1`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 16 | False | 12.945 | 43.705 | 46.859 | 24.180 |
| THUDM/CogVideoX-2b | 16 | True | 12.945 | 12.945 | 21.395 | 24.180 |
| THUDM/CogVideoX-2b | 64 | False | 13.035 | 43.916 | 47.070 | 24.234 |
| THUDM/CogVideoX-2b | 64 | True | 13.035 | 13.035 | 20.887 | 24.266 |
| THUDM/CogVideoX-2b | 256 | False | 13.095 | 44.947 | 48.111 | 24.607 |
| THUDM/CogVideoX-2b | 256 | True | 13.095 | 13.095 | 21.391 | 24.635 |
| THUDM/CogVideoX-5b | 16 | True | 19.742 | 19.742 | 28.533 | 38.002 |
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 20.006 | 29.107 | 38.785 |
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 20.771 | 30.078 | 39.559 |
With `train_batch_size = 4`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 16 | True | 12.945 | 21.709 | 21.762 | 24.254 |
| THUDM/CogVideoX-2b | 64 | True | 13.035 | 21.844 | 21.855 | 24.338 |
| THUDM/CogVideoX-2b | 256 | True | 13.094 | 22.020 | 22.031 | 24.709 |
| THUDM/CogVideoX-5b | 16 | True | 19.742 | 46.262 | 46.297 | 38.400 |
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 46.561 | 46.574 | 38.840 |
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 47.268 | 47.332 | 39.623 |
</details>
<details>
<summary> DeepSpeed (AdamW + CPU/Parameter offloading) </summary>
**Note:** Results are reported with `gradient_checkpointing` enabled, running on a 2x A100.
With `train_batch_size = 1`:
| model | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 13.141 | 13.141 | 21.070 | 24.602 |
| THUDM/CogVideoX-5b | 20.170 | 20.170 | 28.662 | 38.957 |
With `train_batch_size = 4`:
| model | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 13.141 | 19.854 | 20.836 | 24.709 |
| THUDM/CogVideoX-5b | 20.170 | 40.635 | 40.699 | 39.027 |
</details>
### Full finetuning
> [!NOTE]
> The memory requirements for image-to-video full finetuning are similar to that of text-to-video on `THUDM/CogVideoX-5b`, so it hasn't been reported explicitly.
>
> Additionally, to prepare test images for I2V finetuning, you could either generate them on-the-fly by modifying the script, or extract some frames from your training data using:
> `ffmpeg -i input.mp4 -frames:v 1 frame.png`,
> or provide a URL to a valid and accessible image.
> [!NOTE]
> Trying to run full finetuning without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
<details>
<summary> AdamW </summary>
With `train_batch_size = 1`:
| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | True | 16.396 | 33.934 | 43.848 | 37.520 |
| THUDM/CogVideoX-5b | True | 30.061 | OOM | OOM | OOM |
With `train_batch_size = 4`:
| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | True | 16.396 | 38.281 | 48.341 | 37.544 |
| THUDM/CogVideoX-5b | True | 30.061 | OOM | OOM | OOM |
</details>
<details>
<summary> AdamW (8-bit bitsandbytes) </summary>
With `train_batch_size = 1`:
| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | True | 16.396 | 16.447 | 27.555 | 27.156 |
| THUDM/CogVideoX-5b | True | 30.061 | 52.826 | 58.570 | 49.541 |
With `train_batch_size = 4`:
| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | True | 16.396 | 27.930 | 27.990 | 27.326 |
| THUDM/CogVideoX-5b | True | 16.396 | 66.648 | 66.705 | 48.828 |
</details>
<details>
<summary> AdamW + CPUOffloadOptimizer (with gradient offloading) </summary>
With `train_batch_size = 1`:
| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | True | 16.396 | 16.396 | 26.100 | 23.832 |
| THUDM/CogVideoX-5b | True | 30.061 | 39.359 | 48.307 | 37.947 |
With `train_batch_size = 4`:
| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | True | 16.396 | 27.916 | 27.975 | 23.936 |
| THUDM/CogVideoX-5b | True | 30.061 | 66.607 | 66.668 | 38.061 |
</details>
<details>
<summary> DeepSpeed (AdamW + CPU/Parameter offloading) </summary>
**Note:** Results are reported with `gradient_checkpointing` enabled, running on a 2x A100.
With `train_batch_size = 1`:
| model | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 13.111 | 13.111 | 20.328 | 23.867 |
| THUDM/CogVideoX-5b | 19.762 | 19.998 | 27.697 | 38.018 |
With `train_batch_size = 4`:
| model | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
|:------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
| THUDM/CogVideoX-2b | 13.111 | 21.188 | 21.254 | 23.869 |
| THUDM/CogVideoX-5b | 19.762 | 43.465 | 43.531 | 38.082 |
</details>
> [!NOTE]
> - `memory_after_validation` is indicative of the peak memory required for training. This is because apart from the activations, parameters and gradients stored for training, you also need to load the vae and text encoder in memory and spend some memory to perform inference. In order to reduce total memory required to perform training, one can choose not to perform validation/testing as part of the training script.
>
> - `memory_before_validation` is the true indicator of the peak memory required for training if you choose to not perform validation/testing.
<table align="center">
<tr>
<td align="center"><a href="https://www.youtube.com/watch?v=UvRl4ansfCg"> Slaying OOMs with PyTorch</a></td>
</tr>
<tr>
<td align="center"><img src="assets/slaying-ooms.png" style="width: 480px; height: 480px;"></td>
</tr>
</table>
## TODOs
- [x] Make scripts compatible with DDP
- [ ] Make scripts compatible with FSDP
- [x] Make scripts compatible with DeepSpeed
- [ ] vLLM-powered captioning script
- [x] Multi-resolution/frame support in `prepare_dataset.py`
- [ ] Analyzing traces for potential speedups and removing as many syncs as possible
- [x] Test scripts with memory-efficient optimizer from bitsandbytes
- [x] Test scripts with CPUOffloadOptimizer, etc.
- [ ] Test scripts with torchao quantization, and low bit memory optimizers (Currently errors with AdamW (8/4-bit torchao))
- [ ] Test scripts with AdamW (8-bit bitsandbytes) + CPUOffloadOptimizer (with gradient offloading) (Currently errors out)
- [ ] [Sage Attention](https://github.com/thu-ml/SageAttention) (work with the authors to support backward pass, and optimize for A100)
> [!IMPORTANT]
> Since our goal is to make the scripts as memory-friendly as possible we don't guarantee multi-GPU training.