diff --git a/.gitignore b/.gitignore
index bcc6640..ecadaee 100644
--- a/.gitignore
+++ b/.gitignore
@@ -168,5 +168,7 @@ cython_debug/
wandb/
*.txt
dump*
+outputs*
+*.slurm
!requirements.txt
diff --git a/README.md b/README.md
index e9f7bb2..0a60785 100644
--- a/README.md
+++ b/README.md
@@ -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.
@@ -10,8 +10,6 @@ Fine-tune Cog family of video models for custom video generation under 24GB of G
-**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).
+
+
+ LTX Video
+
+### 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("").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_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`:
 |
-
-## 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.
diff --git a/accelerate_configs/uncompiled_8.yaml b/accelerate_configs/uncompiled_8.yaml
new file mode 100644
index 0000000..ee7f50c
--- /dev/null
+++ b/accelerate_configs/uncompiled_8.yaml
@@ -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
\ No newline at end of file
diff --git a/finetrainers/__init__.py b/finetrainers/__init__.py
new file mode 100644
index 0000000..412e298
--- /dev/null
+++ b/finetrainers/__init__.py
@@ -0,0 +1,2 @@
+from .args import Args, parse_arguments
+from .trainer import Trainer
diff --git a/finetrainers/args.py b/finetrainers/args.py
new file mode 100644
index 0000000..a4027ec
--- /dev/null
+++ b/finetrainers/args.py
@@ -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"
diff --git a/finetrainers/constants.py b/finetrainers/constants.py
new file mode 100644
index 0000000..bc050f6
--- /dev/null
+++ b/finetrainers/constants.py
@@ -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
+
+
+
+\#\# 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()
diff --git a/finetrainers/dataset.py b/finetrainers/dataset.py
new file mode 100644
index 0000000..8ca0bb6
--- /dev/null
+++ b/finetrainers/dataset.py
@@ -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] = []
diff --git a/finetrainers/ltx_video/__init__.py b/finetrainers/ltx_video/__init__.py
new file mode 100644
index 0000000..5282217
--- /dev/null
+++ b/finetrainers/ltx_video/__init__.py
@@ -0,0 +1 @@
+from .ltx_video import LTX_VIDEO_T2V_CONFIG
diff --git a/finetrainers/ltx_video/ltx_video.py b/finetrainers/ltx_video/ltx_video.py
new file mode 100644
index 0000000..d1a0559
--- /dev/null
+++ b/finetrainers/ltx_video/ltx_video.py
@@ -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,
+}
diff --git a/finetrainers/models.py b/finetrainers/models.py
new file mode 100644
index 0000000..9e59ec2
--- /dev/null
+++ b/finetrainers/models.py
@@ -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]
diff --git a/finetrainers/state.py b/finetrainers/state.py
new file mode 100644
index 0000000..f30b8c2
--- /dev/null
+++ b/finetrainers/state.py
@@ -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
diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py
new file mode 100644
index 0000000..84107a7
--- /dev/null
+++ b/finetrainers/trainer.py
@@ -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)
diff --git a/finetrainers/utils/__init__.py b/finetrainers/utils/__init__.py
new file mode 100644
index 0000000..bdf27f5
--- /dev/null
+++ b/finetrainers/utils/__init__.py
@@ -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
diff --git a/finetrainers/utils/diffusion_utils.py b/finetrainers/utils/diffusion_utils.py
new file mode 100644
index 0000000..be37bc9
--- /dev/null
+++ b/finetrainers/utils/diffusion_utils.py
@@ -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
diff --git a/finetrainers/utils/file_utils.py b/finetrainers/utils/file_utils.py
new file mode 100644
index 0000000..682563d
--- /dev/null
+++ b/finetrainers/utils/file_utils.py
@@ -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("?", "-")
+ )
diff --git a/finetrainers/utils/memory_utils.py b/finetrainers/utils/memory_utils.py
new file mode 100644
index 0000000..3742d46
--- /dev/null
+++ b/finetrainers/utils/memory_utils.py
@@ -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
diff --git a/finetrainers/utils/optimizer_utils.py b/finetrainers/utils/optimizer_utils.py
new file mode 100644
index 0000000..d05d5d3
--- /dev/null
+++ b/finetrainers/utils/optimizer_utils.py
@@ -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
diff --git a/finetrainers/utils/torch_utils.py b/finetrainers/utils/torch_utils.py
new file mode 100644
index 0000000..32190bb
--- /dev/null
+++ b/finetrainers/utils/torch_utils.py
@@ -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
diff --git a/train.py b/train.py
new file mode 100644
index 0000000..0a45987
--- /dev/null
+++ b/train.py
@@ -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()
diff --git a/train_image_to_video_lora.sh b/train_image_to_video_lora.sh
index c219e8c..8ff0111 100755
--- a/train_image_to_video_lora.sh
+++ b/train_image_to_video_lora.sh
@@ -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 \
diff --git a/train_image_to_video_sft.sh b/train_image_to_video_sft.sh
index 47e8fd7..9497612 100755
--- a/train_image_to_video_sft.sh
+++ b/train_image_to_video_sft.sh
@@ -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 \
diff --git a/train_text_to_video_lora.sh b/train_text_to_video_lora.sh
index ad0bc4a..e7239f5 100755
--- a/train_text_to_video_lora.sh
+++ b/train_text_to_video_lora.sh
@@ -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 \
diff --git a/train_text_to_video_sft.sh b/train_text_to_video_sft.sh
index 9514154..b4de76c 100755
--- a/train_text_to_video_sft.sh
+++ b/train_text_to_video_sft.sh
@@ -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 \
diff --git a/training/README.md b/training/README.md
new file mode 100644
index 0000000..b529d9b
--- /dev/null
+++ b/training/README.md
@@ -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 ⚡️📼
+
+
+
+ |
+
+
+
+**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("").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
+
+
+
+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.
+
+
+ AdamW
+
+**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 |
+
+
+
+
+ AdamW (8-bit bitsandbytes)
+
+**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 |
+
+
+
+
+ AdamW + CPUOffloadOptimizer (with gradient offloading)
+
+**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 |
+
+
+
+
+ DeepSpeed (AdamW + CPU/Parameter offloading)
+
+**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 |
+
+
+
+### 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.
+
+
+ AdamW
+
+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 |
+
+
+
+
+ AdamW (8-bit bitsandbytes)
+
+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 |
+
+
+
+
+ AdamW + CPUOffloadOptimizer (with gradient offloading)
+
+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 |
+
+
+
+
+ DeepSpeed (AdamW + CPU/Parameter offloading)
+
+**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 |
+
+
+
+> [!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.
+
+
+
+## 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.
diff --git a/README_zh.md b/training/README_zh.md
similarity index 100%
rename from README_zh.md
rename to training/README_zh.md
diff --git a/training/__init__.py b/training/cogvideox/__init__.py
similarity index 100%
rename from training/__init__.py
rename to training/cogvideox/__init__.py
diff --git a/training/args.py b/training/cogvideox/args.py
similarity index 100%
rename from training/args.py
rename to training/cogvideox/args.py
diff --git a/training/cogvideox_image_to_video_lora.py b/training/cogvideox/cogvideox_image_to_video_lora.py
similarity index 100%
rename from training/cogvideox_image_to_video_lora.py
rename to training/cogvideox/cogvideox_image_to_video_lora.py
diff --git a/training/cogvideox_image_to_video_sft.py b/training/cogvideox/cogvideox_image_to_video_sft.py
similarity index 100%
rename from training/cogvideox_image_to_video_sft.py
rename to training/cogvideox/cogvideox_image_to_video_sft.py
diff --git a/training/cogvideox_text_to_video_lora.py b/training/cogvideox/cogvideox_text_to_video_lora.py
similarity index 100%
rename from training/cogvideox_text_to_video_lora.py
rename to training/cogvideox/cogvideox_text_to_video_lora.py
diff --git a/training/cogvideox_text_to_video_sft.py b/training/cogvideox/cogvideox_text_to_video_sft.py
similarity index 100%
rename from training/cogvideox_text_to_video_sft.py
rename to training/cogvideox/cogvideox_text_to_video_sft.py
diff --git a/training/dataset.py b/training/cogvideox/dataset.py
similarity index 100%
rename from training/dataset.py
rename to training/cogvideox/dataset.py
diff --git a/training/prepare_dataset.py b/training/cogvideox/prepare_dataset.py
similarity index 100%
rename from training/prepare_dataset.py
rename to training/cogvideox/prepare_dataset.py
diff --git a/training/text_encoder/__init__.py b/training/cogvideox/text_encoder/__init__.py
similarity index 100%
rename from training/text_encoder/__init__.py
rename to training/cogvideox/text_encoder/__init__.py
diff --git a/training/text_encoder/text_encoder.py b/training/cogvideox/text_encoder/text_encoder.py
similarity index 100%
rename from training/text_encoder/text_encoder.py
rename to training/cogvideox/text_encoder/text_encoder.py
diff --git a/training/utils.py b/training/cogvideox/utils.py
similarity index 100%
rename from training/utils.py
rename to training/cogvideox/utils.py