diff --git a/.gitignore b/.gitignore index 82f9275..1a89f04 100644 --- a/.gitignore +++ b/.gitignore @@ -160,3 +160,8 @@ cython_debug/ # and can be added to the global gitignore or merged into this file. For a more nuclear # option (not recommended) you can uncomment the following to ignore the entire idea folder. #.idea/ + +# manually added +wandb/ +*.txt +dump* diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..78644f9 --- /dev/null +++ b/Makefile @@ -0,0 +1,11 @@ +.PHONY: quality style + +check_dirs := training tests + +quality: + ruff check $(check_dirs) + ruff format --check $(check_dirs) setup.py + +style: + ruff check $(check_dirs) --fix + ruff format $(check_dirs) diff --git a/README.md b/README.md index b2fe54e..487d435 100644 --- a/README.md +++ b/README.md @@ -1 +1,101 @@ -# cogvideox-distillation \ No newline at end of file +# Finetuning CogVideoX + + +## Dataset Preparation + +Create two files where one file contains line-separated prompts and another file contains line-separated paths to video data (the path to video files must be relative to the path you pass when specifying `--data_root`). Let's take a look at an example to understand this better! + +Assume you've specified `--data_root` as `/dataset`, and that this directory contains the files: `prompts.txt` and `videos.txt`. + +The `prompts.txt` file should contain line-separated prompts: + +``` +A black and white animated sequence featuring a rabbit, named Rabbity Ribfried, and an anthropomorphic goat in a musical, playful environment, showcasing their evolving interaction. +A black and white animated sequence on a ship's deck features a bulldog character, named Bully Bulldoger, showcasing exaggerated facial expressions and body language. The character progresses from confident to focused, then to strained and distressed, displaying a range of emotions as it navigates challenges. The ship's interior remains static in the background, with minimalistic details such as a bell and open door. The character's dynamic movements and changing expressions drive the narrative, with no camera movement to distract from its evolving reactions and physical gestures. +... +``` + +The `videos.txt` file should contain line-separate paths to video files. Note that the path should be _relative_ to the `--data_root` directory. + +```bash +videos/00000.mp4 +videos/00001.mp4 +... +``` + +Overall, this is how your dataset would look like if you ran the `tree` command on the dataset root directory: + +```bash +/dataset +├── prompts.txt +├── videos.txt +├── videos + ├── videos/00000.mp4 + ├── videos/00001.mp4 + ├── ... +``` + +When using this format, the `--caption_column` must be `prompts.txt` and `--video_column` must be `videos.txt`. If you, instead, have your data stored in a CSV file, you can also specify `--dataset_file` as the path to CSV, the `--caption_column` and `--video_column` as the actual column names in the CSV file. + +As an example, let's use [this](https://huggingface.co/datasets/Wild-Heart/Disney-VideoGeneration-Dataset) Disney dataset for finetuning. To download, one can use the 🤗 Hugging Face CLI. + +```bash +huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset --local-dir video-dataset-disney +``` + +#### Rough notes and TODOs: + +- Uncompiled SFT works end-to-end on dummy example. Need to test on larger dataset (not priority at the moment) +- Compiled SFT fails with `THUDM/CogVideoX-2b` throwing the following error (by error, it's more of a graph break situation due to mixin numpy/cpu device when getting sincos positional embeddings). + +## Training + +TODO + +Take a look at `training/*.sh` + +Note: Untested on MPS + +## Memory requirements + +| model | lora rank | optimizer | gradient_checkpointing | memory_before_training | memory_after_validation | memory_after_testing | +|:------------------:|:---------:|:---------:|:----------------------:|:----------------------:|:-----------------------:|:--------------------:| +| THUDM/CogVideoX-2b | 16 | adamw | False | 12.945 | 39.553 | 23.148 | +| THUDM/CogVideoX-2b | 16 | adamw | True | 12.946 | 18.436 | 23.160 | +| THUDM/CogVideoX-2b | 64 | adamw | False | 13.035 | 40.051 | 23.430 | +| THUDM/CogVideoX-2b | 64 | adamw | True | 13.035 | 18.883 | 23.414 | +| THUDM/CogVideoX-2b | 256 | adamw | False | 13.095 | 42.004 | 24.385 | +| THUDM/CogVideoX-2b | 256 | adamw | True | 13.095 | 19.307 | 24.381 | + +**Note:** `memory_after_validation` is indicative of the peak memory required for training. + +
+ stack trace + +``` +skipping cudagraphs due to skipping cudagraphs due to cpu device (cat_3). Found from : + File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/accelerate/utils/operations.py", line 820, in forward + return model_forward(*args, **kwargs) + File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/accelerate/utils/operations.py", line 808, in __call__ + return convert_to_fp32(self.model_forward(*args, **kwargs)) + File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/amp/autocast_mode.py", line 44, in decorate_autocast + return func(*args, **kwargs) + File "/home/aryan/work/diffusers/src/diffusers/models/transformers/cogvideox_transformer_3d.py", line 446, in forward + hidden_states = self.patch_embed(encoder_hidden_states, hidden_states) + File "/home/aryan/work/diffusers/src/diffusers/models/embeddings.py", line 435, in forward + pos_embedding = self._get_positional_embeddings(height, width, pre_time_compression_frames) + File "/home/aryan/work/diffusers/src/diffusers/models/embeddings.py", line 385, in _get_positional_embeddings + pos_embedding = get_3d_sincos_pos_embed( + File "/home/aryan/work/diffusers/src/diffusers/models/embeddings.py", line 108, in get_3d_sincos_pos_embed + grid = np.stack(grid, axis=0) +``` +
+ +- Make T2V LoRA script up-to-date +- Make I2V LoRA script up-to-date +- Make scripts compatible with DDP +- Make scripts compatible with FSDP +- Make scripts compatible with DeepSpeed +- Test scripts with memory-efficient optimizer +- Test scripts with quantization using torchao, CPUOffloadOptimizer, etc. +- Make 5B lora finetuning work in under 24GB diff --git a/accelerate_configs/compiled_1.yaml b/accelerate_configs/compiled_1.yaml new file mode 100644 index 0000000..646cb6b --- /dev/null +++ b/accelerate_configs/compiled_1.yaml @@ -0,0 +1,22 @@ +compute_environment: LOCAL_MACHINE +debug: false +distributed_type: 'NO' +downcast_bf16: 'no' +dynamo_config: + dynamo_backend: INDUCTOR + dynamo_mode: max-autotune + dynamo_use_dynamic: true + dynamo_use_fullgraph: false +enable_cpu_affinity: false +gpu_ids: '3' +machine_rank: 0 +main_training_function: main +mixed_precision: fp16 +num_machines: 1 +num_processes: 1 +rdzv_backend: static +same_network: true +tpu_env: [] +tpu_use_cluster: false +tpu_use_sudo: false +use_cpu: false diff --git a/accelerate_configs/uncompiled_1.yaml b/accelerate_configs/uncompiled_1.yaml new file mode 100644 index 0000000..f81112a --- /dev/null +++ b/accelerate_configs/uncompiled_1.yaml @@ -0,0 +1,17 @@ +compute_environment: LOCAL_MACHINE +debug: false +distributed_type: 'NO' +downcast_bf16: 'no' +enable_cpu_affinity: false +gpu_ids: '3' +machine_rank: 0 +main_training_function: main +mixed_precision: fp16 +num_machines: 1 +num_processes: 1 +rdzv_backend: static +same_network: true +tpu_env: [] +tpu_use_cluster: false +tpu_use_sudo: false +use_cpu: false diff --git a/assets/tests/metadata.csv b/assets/tests/metadata.csv new file mode 100644 index 0000000..ac6f2df --- /dev/null +++ b/assets/tests/metadata.csv @@ -0,0 +1,2 @@ +video,caption +"videos/hiker.mp4","""A hiker standing at the top of a mountain, triumphantly, high quality""" \ No newline at end of file diff --git a/assets/tests/prompts.txt b/assets/tests/prompts.txt new file mode 100644 index 0000000..866219d --- /dev/null +++ b/assets/tests/prompts.txt @@ -0,0 +1 @@ +A hiker standing at the top of a mountain, triumphantly, high quality \ No newline at end of file diff --git a/assets/tests/prompts_multi.txt b/assets/tests/prompts_multi.txt new file mode 100644 index 0000000..684c3db --- /dev/null +++ b/assets/tests/prompts_multi.txt @@ -0,0 +1,2 @@ +A hiker standing at the top of a mountain, triumphantly, high quality +A hiker standing at the top of a mountain, triumphantly, high quality \ No newline at end of file diff --git a/assets/tests/videos.txt b/assets/tests/videos.txt new file mode 100644 index 0000000..f60f5cd --- /dev/null +++ b/assets/tests/videos.txt @@ -0,0 +1 @@ +videos/hiker.mp4 \ No newline at end of file diff --git a/assets/tests/videos/hiker.mp4 b/assets/tests/videos/hiker.mp4 new file mode 100644 index 0000000..c2e10c9 Binary files /dev/null and b/assets/tests/videos/hiker.mp4 differ diff --git a/assets/tests/videos/hiker_tiny.mp4 b/assets/tests/videos/hiker_tiny.mp4 new file mode 100644 index 0000000..da082ea Binary files /dev/null and b/assets/tests/videos/hiker_tiny.mp4 differ diff --git a/assets/tests/videos_multi.txt b/assets/tests/videos_multi.txt new file mode 100644 index 0000000..66f3c8a --- /dev/null +++ b/assets/tests/videos_multi.txt @@ -0,0 +1,2 @@ +videos/hiker.mp4 +videos/hiker_tiny.mp4 \ No newline at end of file diff --git a/metrics.sh b/metrics.sh new file mode 100755 index 0000000..4341cbe --- /dev/null +++ b/metrics.sh @@ -0,0 +1,155 @@ +# 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 + +# GPU_IDS="1" +# LEARNING_RATES=("1e-4") +# LR_SCHEDULES=("cosine_with_restarts") +# OPTIMIZERS=("adamw") +# MAX_TRAIN_STEPS=("2") +# RANK=("16" "64" "256") +# GRADIENT_CHECKPOINTING=("" "--gradient_checkpointing") + +# DATA_ROOT="/raid/aryan/video-dataset-disney/" +# CAPTION_COLUMN="prompts.txt" +# VIDEO_COLUMN="videos.txt" + +# 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 +# for rank in "${RANK[@]}"; do +# for gradient_checkpointing in "${GRADIENT_CHECKPOINTING[@]}"; do +# cache_dir="/raid/aryan/cogvideox-lora/" +# output_dir="/raid/aryan/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/" + +# cmd="accelerate launch --config_file accelerate_configs/uncompiled_1.yaml --gpu_ids $GPU_IDS training/cogvideox_text_to_video_lora.py \ +# --pretrained_model_name_or_path THUDM/CogVideoX-2b \ +# --cache_dir $cache_dir \ +# --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 \ +# --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 1 \ +# --seed 42 \ +# --rank $rank \ +# --lora_alpha 64 \ +# --mixed_precision fp16 \ +# --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 200 \ +# --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 +# done +# done + +# For testing load from tensor data +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 + +GPU_IDS="1" +LEARNING_RATES=("1e-4") +LR_SCHEDULES=("cosine_with_restarts") +OPTIMIZERS=("adamw") +MAX_TRAIN_STEPS=("2") +RANK=("16" "64" "256") +GRADIENT_CHECKPOINTING=("" "--gradient_checkpointing") + +DATA_ROOT="training/dump" +CAPTION_COLUMN="prompts.txt" +VIDEO_COLUMN="videos.txt" + +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 + for rank in "${RANK[@]}"; do + for gradient_checkpointing in "${GRADIENT_CHECKPOINTING[@]}"; do + cache_dir="/raid/aryan/cogvideox-lora/" + output_dir="/raid/aryan/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/" + + cmd="accelerate launch --config_file accelerate_configs/uncompiled_1.yaml --gpu_ids $GPU_IDS training/cogvideox_text_to_video_lora.py \ + --pretrained_model_name_or_path THUDM/CogVideoX-2b \ + --cache_dir $cache_dir \ + --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 \ + --load_tensors \ + --validation_prompt \"BW_STYLE A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions\" \ + --validation_prompt_separator ::: \ + --num_validation_videos 1 \ + --validation_epochs 2 \ + --seed 42 \ + --rank $rank \ + --lora_alpha 64 \ + --mixed_precision fp16 \ + --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 200 \ + --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 + done +done diff --git a/prepare_dataset.sh b/prepare_dataset.sh new file mode 100755 index 0000000..b0414e9 --- /dev/null +++ b/prepare_dataset.sh @@ -0,0 +1,42 @@ +#!/bin/bash + +MODEL_ID="THUDM/CogVideoX-2b" + +# For more details on the expected data format, please refer to the README. +DATA_ROOT="/raid/aryan/video-dataset-tom-and-jerry" # This needs to be the path to the base directory where your videos are located. +CAPTION_COLUMN="prompts.txt" +VIDEO_COLUMN="videos.txt" +OUTPUT_DIR="/raid/aryan/video-dataset-tom-and-jerry-encoded" +HEIGHT=480 +WIDTH=720 +MAX_NUM_FRAMES=49 +MAX_SEQUENCE_LENGTH=226 +TARGET_FPS=8 +BATCH_SIZE=1 +DTYPE=fp32 + +# To create a folder-style dataset structure without pre-encoding videos and captions' +CMD_WITHOUT_PRE_ENCODING="\ + python3 training/prepare_dataset.py \ + --model_id $MODEL_ID \ + --data_root $DATA_ROOT \ + --caption_column $CAPTION_COLUMN \ + --video_column $VIDEO_COLUMN \ + --output_dir $OUTPUT_DIR \ + --height $HEIGHT \ + --width $WIDTH \ + --max_num_frames $MAX_NUM_FRAMES \ + --max_sequence_length $MAX_SEQUENCE_LENGTH \ + --target_fps $TARGET_FPS \ + --batch_size $BATCH_SIZE \ + --dtype $DTYPE +" + +CMD_WITH_PRE_ENCODING="$CMD_WITHOUT_PRE_ENCODING --save_tensors" + +# Select which you'd like to run +CMD=$CMD_WITH_PRE_ENCODING + +echo "===== Running \`$CMD\` =====" +eval $CMD +echo -ne "===== Finished running script =====\n" diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..79d64d3 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,28 @@ +[tool.ruff] +line-length = 119 + +[tool.ruff.lint] +# Never enforce `E501` (line length violations). +ignore = ["C901", "E501", "E741", "F402", "F823"] +select = ["C", "E", "F", "I", "W"] + +# Ignore import violations in all `__init__.py` files. +[tool.ruff.lint.per-file-ignores] +"__init__.py" = ["E402", "F401", "F403", "F811"] + +[tool.ruff.lint.isort] +lines-after-imports = 2 +known-first-party = [] + +[tool.ruff.format] +# Like Black, use double quotes for strings. +quote-style = "double" + +# Like Black, indent with spaces, rather than tabs. +indent-style = "space" + +# Like Black, respect magic trailing commas. +skip-magic-trailing-comma = false + +# Like Black, automatically detect the appropriate line ending. +line-ending = "auto" diff --git a/tests/test_dataset.py b/tests/test_dataset.py new file mode 100644 index 0000000..344cb6b --- /dev/null +++ b/tests/test_dataset.py @@ -0,0 +1,104 @@ +# Run: python3 tests/test_dataset.py + +import sys + + +def test_video_dataset(): + from dataset import VideoDataset + + dataset_dirs = VideoDataset( + data_root="assets/tests/", + caption_column="prompts.txt", + video_column="videos.txt", + max_num_frames=49, + id_token=None, + random_flip=None, + ) + dataset_csv = VideoDataset( + data_root="assets/tests/", + dataset_file="assets/tests/metadata.csv", + caption_column="caption", + video_column="video", + max_num_frames=49, + id_token=None, + random_flip=None, + ) + + assert len(dataset_dirs) == 1 + assert len(dataset_csv) == 1 + assert dataset_dirs[0]["video"].shape == (49, 3, 480, 720) + assert (dataset_dirs[0]["video"] == dataset_csv[0]["video"]).all() + + print(dataset_dirs[0]["video"].shape) + + +def test_video_dataset_with_resizing(): + from dataset import VideoDatasetWithResizing + + dataset_dirs = VideoDatasetWithResizing( + data_root="assets/tests/", + caption_column="prompts.txt", + video_column="videos.txt", + max_num_frames=49, + id_token=None, + random_flip=None, + ) + dataset_csv = VideoDatasetWithResizing( + data_root="assets/tests/", + dataset_file="assets/tests/metadata.csv", + caption_column="caption", + video_column="video", + max_num_frames=49, + id_token=None, + random_flip=None, + ) + + assert len(dataset_dirs) == 1 + assert len(dataset_csv) == 1 + assert dataset_dirs[0]["video"].shape == (48, 3, 480, 720) # Changes due to T2V frame bucket sampling + assert (dataset_dirs[0]["video"] == dataset_csv[0]["video"]).all() + + print(dataset_dirs[0]["video"].shape) + + +def test_video_dataset_with_bucket_sampler(): + import torch + from dataset import BucketSampler, VideoDatasetWithResizing + from torch.utils.data import DataLoader + + dataset_dirs = VideoDatasetWithResizing( + data_root="assets/tests/", + caption_column="prompts_multi.txt", + video_column="videos_multi.txt", + max_num_frames=49, + id_token=None, + random_flip=None, + ) + sampler = BucketSampler(dataset_dirs, batch_size=8) + + def collate_fn(data): + captions = [x["prompt"] for x in data[0]] + videos = [x["video"] for x in data[0]] + videos = torch.stack(videos) + return captions, videos + + dataloader = DataLoader(dataset_dirs, batch_size=1, sampler=sampler, collate_fn=collate_fn) + first = False + + for captions, videos in dataloader: + if not first: + assert len(captions) == 8 and isinstance(captions[0], str) + assert videos.shape == (8, 48, 3, 480, 720) + first = True + else: + assert len(captions) == 8 and isinstance(captions[0], str) + assert videos.shape == (8, 48, 3, 256, 360) + break + + +if __name__ == "__main__": + sys.path.append("./training") + + test_video_dataset() + test_video_dataset_with_resizing() + test_video_dataset_with_bucket_sampler() diff --git a/train_text_to_video_lora.sh b/train_text_to_video_lora.sh new file mode 100755 index 0000000..6e9931b --- /dev/null +++ b/train_text_to_video_lora.sh @@ -0,0 +1,70 @@ +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 + +GPU_IDS="2" +LEARNING_RATES=("1e-4") +LR_SCHEDULES=("cosine_with_restarts") +OPTIMIZERS=("adamw") +MAX_TRAIN_STEPS=("2") + +DATA_ROOT="dump" +CAPTION_COLUMN="prompts.txt" +VIDEO_COLUMN="videos.txt" + +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 + cache_dir="/raid/aryan/cogvideox-lora/" + output_dir="/raid/aryan/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/" + + cmd="accelerate launch --config_file accelerate_configs/uncompiled_1.yaml --gpu_ids $GPU_IDS training/cogvideox_text_to_video_lora.py \ + --pretrained_model_name_or_path THUDM/CogVideoX-2b \ + --cache_dir $cache_dir \ + --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 \ + --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 64 \ + --lora_alpha 64 \ + --mixed_precision fp16 \ + --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 200 \ + --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 diff --git a/train_text_to_video_sft.sh b/train_text_to_video_sft.sh new file mode 100755 index 0000000..22dabc3 --- /dev/null +++ b/train_text_to_video_sft.sh @@ -0,0 +1,68 @@ +# 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 + +GPU_IDS="3" +LEARNING_RATES=("1e-4") +LR_SCHEDULES=("cosine_with_restarts") +OPTIMIZERS=("adamw") +MAX_TRAIN_STEPS=("20000") + +# DATA_ROOT="/raid/aryan/dataset-cogvideox/" +DATA_ROOT="/raid/aryan/openvid-1m" +CAPTION_COLUMN="prompts.txt" +VIDEO_COLUMN="videos.txt" + +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 + cache_dir="/raid/aryan/cogvideox-sft/" + output_dir="/raid/aryan/cogvideox-sft__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/" + + cmd="accelerate launch --config_file accelerate_configs/uncompiled_1.yaml --gpu_ids $GPU_IDS training/cogvideox_text_to_video_sft.py \ + --pretrained_model_name_or_path THUDM/CogVideoX-2b \ + --cache_dir $cache_dir \ + --data_root $DATA_ROOT \ + --caption_column $CAPTION_COLUMN \ + --video_column $VIDEO_COLUMN \ + --height_buckets 480 \ + --width_buckets 720 \ + --frame_buckets 49 \ + --validation_prompt \"a man wearing a bicycle helmet, riding a bike through a forested area. The man is wearing a black t-shirt and appears to be in motion, as suggested by the slight blur of the background. The forest is lush and green, with trees and foliage filling the background. The man's helmet is white with a black visor, and he is looking directly at the camera with a slight smile on his face. The style of the video is casual and candid, capturing a moment of outdoor activity:::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 1 \ + --seed 42 \ + --mixed_precision fp16 \ + --output_dir $output_dir \ + --max_num_frames 49 \ + --train_batch_size 1 \ + --max_train_steps $steps \ + --checkpointing_steps 2000 \ + --gradient_accumulation_steps 1 \ + --gradient_checkpointing \ + --learning_rate $learning_rate \ + --lr_scheduler $lr_schedule \ + --lr_warmup_steps 200 \ + --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 diff --git a/training/args.py b/training/args.py new file mode 100644 index 0000000..fa8d6d4 --- /dev/null +++ b/training/args.py @@ -0,0 +1,420 @@ +import argparse + + +def _get_model_args(parser: argparse.ArgumentParser) -> None: + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + + +def _get_dataset_args(parser: argparse.ArgumentParser) -> None: + parser.add_argument( + "--data_root", + type=str, + default=None, + help=("A folder containing the training data."), + ) + parser.add_argument( + "--dataset_file", + type=str, + default=None, + help=("Path to a CSV file if loading prompts/video paths using this format."), + ) + parser.add_argument( + "--video_column", + type=str, + default="video", + help="The column of the dataset containing videos. Or, the name of the file in `--data_root` folder containing the line-separated path to video data.", + ) + parser.add_argument( + "--caption_column", + type=str, + default="text", + help="The column of the dataset containing the instance prompt for each video. Or, the name of the file in `--data_root` folder containing the line-separated instance prompts.", + ) + parser.add_argument( + "--id_token", + type=str, + default=None, + help="Identifier token appended to the start of each prompt if provided.", + ) + parser.add_argument( + "--height_buckets", + nargs="+", + type=int, + default=[256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536], + ) + parser.add_argument( + "--width_buckets", + nargs="+", + type=int, + default=[256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536], + ) + parser.add_argument( + "--frame_buckets", + nargs="+", + type=int, + default=[49], + ) + parser.add_argument( + "--load_tensors", + action="store_true", + help="Whether to use a pre-encoded tensor dataset of latents and prompt embeddings instead of videos and text prompts. The expected format is that saved by running the `prepare_dataset.py` script.", + ) + parser.add_argument( + "--random_flip", + type=float, + default=None, + help="If random horizontal flip augmentation is to be used, this should be the flip probability.", + ) + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.", + ) + + +def _get_validation_args(parser: argparse.ArgumentParser) -> None: + parser.add_argument( + "--validation_prompt", + type=str, + default=None, + help="One or more prompt(s) that is used during validation to verify that the model is learning. Multiple validation prompts should be separated by the '--validation_prompt_seperator' string.", + ) + parser.add_argument( + "--validation_prompt_separator", + type=str, + default=":::", + help="String that separates multiple validation prompts", + ) + parser.add_argument( + "--num_validation_videos", + type=int, + default=1, + help="Number of videos that should be generated during validation per `validation_prompt`.", + ) + parser.add_argument( + "--validation_epochs", + type=int, + default=50, + help="Run validation every X training steps. Validation consists of running the validation prompt `args.num_validation_videos` times.", + ) + parser.add_argument( + "--guidance_scale", + type=float, + default=6, + help="The guidance scale to use while sampling validation videos.", + ) + parser.add_argument( + "--use_dynamic_cfg", + action="store_true", + default=False, + help="Whether or not to use the default cosine dynamic guidance schedule when sampling validation videos.", + ) + + +def _get_training_args(parser: argparse.ArgumentParser) -> None: + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument("--rank", type=int, default=64, help="The rank for LoRA matrices.") + parser.add_argument( + "--lora_alpha", + type=int, + default=64, + help="The lora_alpha to compute scaling factor (lora_alpha / rank) for LoRA matrices.", + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >= 1.10.and an Nvidia Ampere GPU. " + "Default to the value of accelerate config of the current system or the flag passed with the `accelerate.launch` command. Use this " + "argument to override the accelerate config." + ), + ) + parser.add_argument( + "--output_dir", + type=str, + default="cogvideox-sft", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument( + "--height", + type=int, + default=480, + help="All input videos are resized to this height.", + ) + parser.add_argument( + "--width", + type=int, + default=720, + help="All input videos are resized to this width.", + ) + parser.add_argument("--fps", type=int, default=8, help="All input videos will be used at this FPS.") + parser.add_argument( + "--max_num_frames", + type=int, + default=49, + help="All input videos will be truncated to these many frames.", + ) + parser.add_argument( + "--skip_frames_start", + type=int, + default=0, + help="Number of frames to skip from the beginning of each input video. Useful if training data contains intro sequences.", + ) + parser.add_argument( + "--skip_frames_end", + type=int, + default=0, + help="Number of frames to skip from the end of each input video. Useful if training data contains outro sequences.", + ) + parser.add_argument( + "--train_batch_size", + type=int, + default=4, + help="Batch size (per device) for the training dataloader.", + ) + parser.add_argument("--num_train_epochs", type=int, default=1) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides `--num_train_epochs`.", + ) + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints can be used both as final" + " checkpoints in case they are better than the last checkpoint, and are also suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", + type=int, + default=500, + help="Number of steps for the warmup in the lr scheduler.", + ) + parser.add_argument( + "--lr_num_cycles", + type=int, + default=1, + help="Number of hard resets of the lr in cosine_with_restarts scheduler.", + ) + parser.add_argument( + "--lr_power", + type=float, + default=1.0, + help="Power factor of the polynomial scheduler.", + ) + parser.add_argument( + "--enable_slicing", + action="store_true", + default=False, + help="Whether or not to use VAE slicing for saving memory.", + ) + parser.add_argument( + "--enable_tiling", + action="store_true", + default=False, + help="Whether or not to use VAE tiling for saving memory.", + ) + + +def _get_optimizer_args(parser: argparse.ArgumentParser) -> None: + parser.add_argument( + "--optimizer", + type=lambda s: s.lower(), + default="adam", + choices=["adam", "adamw", "prodigy"], + help=("The optimizer type to use."), + ) + parser.add_argument( + "--use_8bit", + action="store_true", + help="Whether or not to use 8-bit optimizers from `bitsandbytes`. Ignored if incompatible optimzer selected.", + ) + parser.add_argument( + "--beta1", + type=float, + default=0.9, + help="The beta1 parameter for the Adam and Prodigy optimizers.", + ) + parser.add_argument( + "--beta2", + type=float, + default=0.95, + help="The beta2 parameter for the Adam and Prodigy optimizers.", + ) + parser.add_argument( + "--beta3", + type=float, + default=None, + help="Coefficients for computing the Prodigy optimizer's stepsize using running averages. If set to None, uses the value of square root of beta2.", + ) + parser.add_argument( + "--prodigy_decouple", + action="store_true", + help="Use AdamW style decoupled weight decay.", + ) + parser.add_argument( + "--weight_decay", + type=float, + default=1e-04, + help="Weight decay to use for optimizer.", + ) + parser.add_argument( + "--epsilon", + type=float, + default=1e-8, + help="Epsilon value for the Adam optimizer and Prodigy optimizers.", + ) + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument( + "--prodigy_use_bias_correction", + action="store_true", + help="Turn on Adam's bias correction.", + ) + parser.add_argument( + "--prodigy_safeguard_warmup", + action="store_true", + help="Remove lr from the denominator of D estimate to avoid issues during warm-up stage.", + ) + + +def _get_configuration_args(parser: argparse.ArgumentParser) -> None: + parser.add_argument("--tracker_name", type=str, default=None, help="Project tracker name") + parser.add_argument( + "--push_to_hub", + action="store_true", + help="Whether or not to push the model to the Hub.", + ) + parser.add_argument( + "--hub_token", + type=str, + default=None, + help="The token to use to push to the Model Hub.", + ) + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help="Directory where logs are stored.", + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument( + "--nccl_timeout", + type=int, + default=600, + help="Maximum timeout duration before which allgather, or related, operations fail in multi-GPU/multi-node training settings.", + ) + parser.add_argument( + "--report_to", + type=str, + default=None, + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + + +def get_args(): + parser = argparse.ArgumentParser(description="Simple example of a training script for CogVideoX.") + + _get_model_args(parser) + _get_dataset_args(parser) + _get_training_args(parser) + _get_validation_args(parser) + _get_optimizer_args(parser) + _get_configuration_args(parser) + + return parser.parse_args() diff --git a/training/cogvideox_text_to_video_lora.py b/training/cogvideox_text_to_video_lora.py new file mode 100644 index 0000000..9d52149 --- /dev/null +++ b/training/cogvideox_text_to_video_lora.py @@ -0,0 +1,910 @@ +# Copyright 2024 The HuggingFace Team. +# All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import gc +import logging +import math +import os +import shutil +from datetime import timedelta +from pathlib import Path +from typing import Any, Dict + +import diffusers +import torch +import transformers +import wandb +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.utils import ( + DistributedDataParallelKwargs, + InitProcessGroupKwargs, + ProjectConfiguration, + set_seed, +) +from diffusers import ( + AutoencoderKLCogVideoX, + CogVideoXDPMScheduler, + CogVideoXPipeline, + CogVideoXTransformer3DModel, +) +from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution +from diffusers.optimization import get_scheduler +from diffusers.training_utils import cast_training_params +from diffusers.utils import ( + convert_unet_state_dict_to_peft, + export_to_video, + is_wandb_available, +) +from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card +from diffusers.utils.torch_utils import is_compiled_module +from huggingface_hub import create_repo, upload_folder +from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict +from torch.utils.data import DataLoader +from tqdm.auto import tqdm +from transformers import AutoTokenizer, T5EncoderModel + + +from args import get_args # isort:skip +from dataset import BucketSampler, VideoDatasetWithResizing # isort:skip +from text_encoder import compute_prompt_embeddings # isort:skip +from utils import get_gradient_norm, get_optimizer, prepare_rotary_positional_embeddings, print_memory, reset_memory # isort:skip + + +logger = get_logger(__name__) + + +def save_model_card( + repo_id: str, + videos=None, + base_model: str = None, + validation_prompt=None, + repo_folder=None, + fps=8, +): + widget_dict = [] + if videos is not None: + for i, video in enumerate(videos): + export_to_video(video, os.path.join(repo_folder, f"final_video_{i}.mp4", fps=fps)) + widget_dict.append( + { + "text": validation_prompt if validation_prompt else " ", + "output": {"url": f"video_{i}.mp4"}, + } + ) + + model_description = f""" +# CogVideoX LoRA - {repo_id} + + + +## Model description + +These are {repo_id} LoRA weights for {base_model}. + +The weights were trained using the [CogVideoX Diffusers trainer](https://github.com/huggingface/diffusers/blob/main/examples/cogvideo/train_cogvideox_lora.py). + +Was LoRA for the text encoder enabled? No. + +## Download model + +[Download the *.safetensors LoRA]({repo_id}/tree/main) in the Files & versions tab. + +## Use it with the [🧨 diffusers library](https://github.com/huggingface/diffusers) + +```py +from diffusers import CogVideoXPipeline +import torch + +pipe = CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16).to("cuda") +pipe.load_lora_weights("{repo_id}", weight_name="pytorch_lora_weights.safetensors", adapter_name=["cogvideox-lora"]) + +# The LoRA adapter weights are determined by what was used for training. +# In this case, we assume `--lora_alpha` is 32 and `--rank` is 64. +# It can be made lower or higher from what was used in training to decrease or amplify the effect +# of the LoRA upto a tolerance, beyond which one might notice no effect at all or overflows. +pipe.set_adapters(["cogvideox-lora"], [32 / 64]) + +video = pipe("{validation_prompt}", guidance_scale=6, use_dynamic_cfg=True).frames[0] +``` + +For more details, including weighting, merging and fusing LoRAs, check the [documentation on loading LoRAs in diffusers](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters) + +## License + +Please adhere to the licensing terms as described [here](https://huggingface.co/THUDM/CogVideoX-5b/blob/main/LICENSE) and [here](https://huggingface.co/THUDM/CogVideoX-2b/blob/main/LICENSE). +""" + model_card = load_or_create_model_card( + repo_id_or_path=repo_id, + from_training=True, + license="other", + base_model=base_model, + prompt=validation_prompt, + model_description=model_description, + widget=widget_dict, + ) + tags = [ + "text-to-video", + "diffusers-training", + "diffusers", + "lora", + "cogvideox", + "cogvideox-diffusers", + "template:sd-lora", + ] + + model_card = populate_model_card(model_card, tags=tags) + model_card.save(os.path.join(repo_folder, "README.md")) + + +def log_validation( + accelerator: Accelerator, + pipe: CogVideoXPipeline, + args: Dict[str, Any], + pipeline_args: Dict[str, Any], + epoch, + is_final_validation: bool = False, +): + logger.info( + f"Running validation... \n Generating {args.num_validation_videos} videos with prompt: {pipeline_args['prompt']}." + ) + + pipe = pipe.to(accelerator.device) + + # run inference + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) if args.seed else None + + videos = [] + for _ in range(args.num_validation_videos): + video = pipe(**pipeline_args, generator=generator, output_type="np").frames[0] + videos.append(video) + + for tracker in accelerator.trackers: + phase_name = "test" if is_final_validation else "validation" + if tracker.name == "wandb": + video_filenames = [] + for i, video in enumerate(videos): + prompt = ( + pipeline_args["prompt"][:25] + .replace(" ", "_") + .replace(" ", "_") + .replace("'", "_") + .replace('"', "_") + .replace("/", "_") + ) + filename = os.path.join(args.output_dir, f"{phase_name}_video_{i}_{prompt}.mp4") + export_to_video(video, filename, fps=8) + video_filenames.append(filename) + + tracker.log( + { + phase_name: [ + wandb.Video(filename, caption=f"{i}: {pipeline_args['prompt']}") + for i, filename in enumerate(video_filenames) + ] + } + ) + + return videos + + +def main(args): + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if torch.backends.mps.is_available() and args.mixed_precision == "bf16": + # due to pytorch#99272, MPS does not yet support bfloat16. + raise ValueError( + "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." + ) + + logging_dir = Path(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + init_process_group_kwargs = InitProcessGroupKwargs(backend="nccl", timeout=timedelta(seconds=args.nccl_timeout)) + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + kwargs_handlers=[ddp_kwargs, init_process_group_kwargs], + ) + + # Disable AMP for MPS. + if torch.backends.mps.is_available(): + accelerator.native_amp = False + + if args.report_to == "wandb": + if not is_wandb_available(): + raise ImportError("Make sure to install wandb if you want to use it for logging during training.") + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + if args.push_to_hub: + repo_id = create_repo( + repo_id=args.hub_model_id or Path(args.output_dir).name, + exist_ok=True, + ).repo_id + + # Prepare models and scheduler + tokenizer = AutoTokenizer.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="tokenizer", + revision=args.revision, + ) + + text_encoder = T5EncoderModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="text_encoder", + revision=args.revision, + ) + + # CogVideoX-2b weights are stored in float16 + # CogVideoX-5b and CogVideoX-5b-I2V weights are stored in bfloat16 + load_dtype = torch.bfloat16 if "5b" in args.pretrained_model_name_or_path.lower() else torch.float16 + transformer = CogVideoXTransformer3DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=load_dtype, + revision=args.revision, + variant=args.variant, + ) + + vae = AutoencoderKLCogVideoX.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae", + revision=args.revision, + variant=args.variant, + ) + + scheduler = CogVideoXDPMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + + if args.enable_slicing: + vae.enable_slicing() + if args.enable_tiling: + vae.enable_tiling() + + # We only train the additional adapter LoRA layers + text_encoder.requires_grad_(False) + transformer.requires_grad_(False) + vae.requires_grad_(False) + + VAE_SCALING_FACTOR = vae.config.scaling_factor + VAE_SCALE_FACTOR_SPATIAL = 2 ** (len(vae.config.block_out_channels) - 1) + + # For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.state.deepspeed_plugin: + # DeepSpeed is handling precision, use what's in the DeepSpeed config + if ( + "fp16" in accelerator.state.deepspeed_plugin.deepspeed_config + and accelerator.state.deepspeed_plugin.deepspeed_config["fp16"]["enabled"] + ): + weight_dtype = torch.float16 + if ( + "bf16" in accelerator.state.deepspeed_plugin.deepspeed_config + and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"] + ): + weight_dtype = torch.float16 + else: + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + + if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: + # due to pytorch#99272, MPS does not yet support bfloat16. + raise ValueError( + "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." + ) + + text_encoder.to(accelerator.device, dtype=weight_dtype) + transformer.to(accelerator.device, dtype=weight_dtype) + vae.to(accelerator.device, dtype=weight_dtype) + + if args.gradient_checkpointing: + transformer.enable_gradient_checkpointing() + + # now we will add new LoRA weights to the attention layers + transformer_lora_config = LoraConfig( + r=args.rank, + lora_alpha=args.lora_alpha, + init_lora_weights=True, + target_modules=["to_k", "to_q", "to_v", "to_out.0"], + ) + transformer.add_adapter(transformer_lora_config) + + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + transformer_lora_layers_to_save = None + + for model in models: + if isinstance(model, type(unwrap_model(transformer))): + 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 + weights.pop() + + CogVideoXPipeline.save_lora_weights( + output_dir, + transformer_lora_layers=transformer_lora_layers_to_save, + ) + + def load_model_hook(models, input_dir): + transformer_ = None + + while len(models) > 0: + model = models.pop() + + if isinstance(model, type(unwrap_model(transformer))): + transformer_ = model + else: + raise ValueError(f"Unexpected save model: {model.__class__}") + + lora_state_dict = CogVideoXPipeline.lora_state_dict(input_dir) + + transformer_state_dict = { + f'{k.replace("transformer.", "")}': v for k, v in lora_state_dict.items() if k.startswith("transformer.") + } + transformer_state_dict = convert_unet_state_dict_to_peft(transformer_state_dict) + incompatible_keys = set_peft_model_state_dict(transformer_, transformer_state_dict, adapter_name="default") + if incompatible_keys is not None: + # check only for unexpected keys + unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None) + if unexpected_keys: + logger.warning( + f"Loading adapter weights from state_dict led to unexpected keys not found in the model: " + f" {unexpected_keys}. " + ) + + # Make sure the trainable params are in float32. This is again needed since the base models + # are in `weight_dtype`. More details: + # https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804 + if args.mixed_precision == "fp16": + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params([transformer_]) + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32 and torch.cuda.is_available(): + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Make sure the trainable params are in float32. + if args.mixed_precision == "fp16": + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params([transformer], dtype=torch.float32) + + transformer_lora_parameters = list(filter(lambda p: p.requires_grad, transformer.parameters())) + + # Optimization parameters + transformer_parameters_with_lr = { + "params": transformer_lora_parameters, + "lr": args.learning_rate, + } + params_to_optimize = [transformer_parameters_with_lr] + + use_deepspeed_optimizer = ( + accelerator.state.deepspeed_plugin is not None + and "optimizer" in accelerator.state.deepspeed_plugin.deepspeed_config + ) + use_deepspeed_scheduler = ( + accelerator.state.deepspeed_plugin is not None + and "scheduler" not in accelerator.state.deepspeed_plugin.deepspeed_config + ) + + optimizer = get_optimizer( + params_to_optimize=params_to_optimize, + optimizer_name=args.optimizer, + learning_rate=args.learning_rate, + beta1=args.beta1, + beta2=args.beta2, + beta3=args.beta3, + epsilon=args.epsilon, + weight_decay=args.weight_decay, + prodigy_decouple=args.prodigy_decouple, + prodigy_use_bias_correction=args.prodigy_use_bias_correction, + prodigy_safeguard_warmup=args.prodigy_safeguard_warmup, + use_8bit=args.use_8bit, + use_deepspeed=use_deepspeed_optimizer, + ) + + # Dataset and DataLoader + train_dataset = VideoDatasetWithResizing( + data_root=args.data_root, + dataset_file=args.dataset_file, + caption_column=args.caption_column, + video_column=args.video_column, + max_num_frames=args.max_num_frames, + id_token=args.id_token, + height_buckets=args.height_buckets, + width_buckets=args.width_buckets, + frame_buckets=args.frame_buckets, + load_tensors=args.load_tensors, + random_flip=args.random_flip, + ) + + def collate_fn_without_pre_encoding(data): + prompts = [x["prompt"] for x in data[0]] + + videos = [x["video"] for x in data[0]] + videos = torch.stack(videos) + videos = videos.to(accelerator.device, dtype=weight_dtype) + videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W] + latent_dist = vae.encode(videos).latent_dist + videos = latent_dist.sample() * VAE_SCALING_FACTOR + videos = videos.permute(0, 2, 1, 3, 4) # [B, F, C, H, W] + videos = videos.to(memory_format=torch.contiguous_format).float() + + return { + "videos": videos, + "prompts": prompts, + } + + def collate_fn_with_pre_encoding(data): + prompts = [x["prompt"] for x in data[0]] + prompts = torch.stack(prompts).to(accelerator.device, dtype=weight_dtype) + + videos = [x["video"] for x in data[0]] + videos = torch.stack(videos).to(accelerator.device, dtype=weight_dtype) + videos = DiagonalGaussianDistribution(videos).sample() * VAE_SCALING_FACTOR + videos = videos.permute(0, 2, 1, 3, 4) # [B, F, C, H, W] + videos = videos.to(memory_format=torch.contiguous_format).float() + + return { + "videos": videos, + "prompts": prompts, + } + + train_dataloader = DataLoader( + train_dataset, + batch_size=1, + sampler=BucketSampler(train_dataset, batch_size=args.train_batch_size, shuffle=True), + collate_fn=collate_fn_with_pre_encoding if args.load_tensors else collate_fn_without_pre_encoding, + num_workers=args.dataloader_num_workers, + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + if use_deepspeed_scheduler: + from accelerate.utils import DummyScheduler + + lr_scheduler = DummyScheduler( + name=args.lr_scheduler, + optimizer=optimizer, + total_num_steps=args.max_train_steps * accelerator.num_processes, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + ) + else: + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + num_cycles=args.lr_num_cycles, + power=args.lr_power, + ) + + # Prepare everything with our `accelerator`. + transformer, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer, optimizer, train_dataloader, lr_scheduler + ) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_name = args.tracker_name or "cogvideox-lora" + accelerator.init_trackers(tracker_name, config=vars(args)) + + accelerator.print("===== Memory before training =====") + reset_memory(accelerator.device) + print_memory(accelerator.device) + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + num_trainable_parameters = sum(param.numel() for model in params_to_optimize for param in model["params"]) + + accelerator.print("***** Running training *****") + accelerator.print(f" Num trainable parameters = {num_trainable_parameters}") + accelerator.print(f" Num examples = {len(train_dataset)}") + accelerator.print(f" Num epochs = {args.num_train_epochs}") + accelerator.print(f" Instantaneous batch size per device = {args.train_batch_size}") + accelerator.print(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + accelerator.print(f" Gradient accumulation steps = {args.gradient_accumulation_steps}") + accelerator.print(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if not args.resume_from_checkpoint: + initial_global_step = 0 + else: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the mos recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + first_epoch = global_step // num_update_steps_per_epoch + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + # For DeepSpeed training + model_config = transformer.module.config if hasattr(transformer, "module") else transformer.config + + if args.load_tensors: + del vae, text_encoder + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize(accelerator.device) + + for epoch in range(first_epoch, args.num_train_epochs): + transformer.train() + + for step, batch in enumerate(train_dataloader): + models_to_accumulate = [transformer] + + with accelerator.accumulate(models_to_accumulate): + model_input = batch["videos"] + prompts = batch["prompts"] + + # Encode prompts + if not args.load_tensors: + prompt_embeds = compute_prompt_embeddings( + tokenizer, + text_encoder, + prompts, + model_config.max_text_seq_length, + accelerator.device, + weight_dtype, + requires_grad=False, + ) + else: + prompt_embeds = prompts + + # Sample noise that will be added to the latents + noise = torch.randn_like(model_input) + batch_size, num_frames, num_channels, height, width = model_input.shape + + # Sample a random timestep for each image + timesteps = torch.randint( + 0, + scheduler.config.num_train_timesteps, + (batch_size,), + dtype=torch.int64, + device=model_input.device, + ) + + # Prepare rotary embeds + image_rotary_emb = ( + prepare_rotary_positional_embeddings( + height=height * VAE_SCALE_FACTOR_SPATIAL, + width=width * VAE_SCALE_FACTOR_SPATIAL, + num_frames=num_frames, + vae_scale_factor_spatial=VAE_SCALE_FACTOR_SPATIAL, + patch_size=model_config.patch_size, + attention_head_dim=model_config.attention_head_dim, + device=accelerator.device, + ) + if model_config.use_rotary_positional_embeddings + else None + ) + + # Add noise to the model input according to the noise magnitude at each timestep + # (this is the forward diffusion process) + noisy_model_input = scheduler.add_noise(model_input, noise, timesteps) + + # Predict the noise residual + model_output = transformer( + hidden_states=noisy_model_input, + encoder_hidden_states=prompt_embeds, + timestep=timesteps, + image_rotary_emb=image_rotary_emb, + return_dict=False, + )[0] + + model_pred = scheduler.get_velocity(model_output, noisy_model_input, timesteps) + + alphas_cumprod = scheduler.alphas_cumprod[timesteps] + weights = 1 / (1 - alphas_cumprod) + while len(weights.shape) < len(model_pred.shape): + weights = weights.unsqueeze(-1) + + target = model_input + + loss = torch.mean( + (weights * (model_pred - target) ** 2).reshape(batch_size, -1), + dim=1, + ) + loss = loss.mean() + accelerator.backward(loss) + + if accelerator.sync_gradients: + gradient_norm_before_clip = get_gradient_norm(transformer.parameters()) + accelerator.clip_grad_norm_(transformer.parameters(), args.max_grad_norm) + gradient_norm_after_clip = get_gradient_norm(transformer.parameters()) + + if accelerator.state.deepspeed_plugin is None: + optimizer.step() + optimizer.zero_grad() + + lr_scheduler.step() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + global_step += 1 + + if accelerator.is_main_process: + if global_step % args.checkpointing_steps == 0: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"Removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + logs = { + "loss": loss.detach().item(), + "lr": lr_scheduler.get_last_lr()[0], + "gradient_norm_before_clip": gradient_norm_before_clip, + "gradient_norm_after_clip": gradient_norm_after_clip, + } + progress_bar.set_postfix(**logs) + accelerator.log(logs, step=global_step) + + if global_step >= args.max_train_steps: + break + + if accelerator.is_main_process: + if args.validation_prompt is not None and (epoch + 1) % args.validation_epochs == 0: + accelerator.print("===== Memory before validation =====") + print_memory(accelerator.device) + torch.cuda.synchronize(accelerator.device) + + pipe = CogVideoXPipeline.from_pretrained( + args.pretrained_model_name_or_path, + transformer=unwrap_model(transformer), + scheduler=scheduler, + revision=args.revision, + variant=args.variant, + torch_dtype=weight_dtype, + ) + + if args.enable_slicing: + pipe.vae.enable_slicing() + if args.enable_tiling: + pipe.vae.enable_tiling() + + validation_prompts = args.validation_prompt.split(args.validation_prompt_separator) + for validation_prompt in validation_prompts: + pipeline_args = { + "prompt": validation_prompt, + "guidance_scale": args.guidance_scale, + "use_dynamic_cfg": args.use_dynamic_cfg, + "height": args.height, + "width": args.width, + } + + log_validation( + pipe=pipe, + args=args, + accelerator=accelerator, + pipeline_args=pipeline_args, + epoch=epoch, + ) + + accelerator.print("===== Memory after validation =====") + print_memory(accelerator.device) + reset_memory(accelerator.device) + + del pipe + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize(accelerator.device) + + accelerator.wait_for_everyone() + + if accelerator.is_main_process: + transformer = unwrap_model(transformer) + dtype = ( + torch.float16 + if args.mixed_precision == "fp16" + else torch.bfloat16 + if args.mixed_precision == "bf16" + else torch.float32 + ) + transformer = transformer.to(dtype) + transformer_lora_layers = get_peft_model_state_dict(transformer) + + CogVideoXPipeline.save_lora_weights( + save_directory=args.output_dir, + transformer_lora_layers=transformer_lora_layers, + ) + + # Cleanup trained models to save memory + if args.load_tensors: + del transformer + else: + del transformer, text_encoder, vae + + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize(accelerator.device) + + accelerator.print("===== Memory before testing =====") + print_memory(accelerator.device) + reset_memory(accelerator.device) + + # Final test inference + pipe = CogVideoXPipeline.from_pretrained( + args.pretrained_model_name_or_path, + revision=args.revision, + variant=args.variant, + torch_dtype=weight_dtype, + ) + pipe.scheduler = CogVideoXDPMScheduler.from_config(pipe.scheduler.config) + + if args.enable_slicing: + pipe.vae.enable_slicing() + if args.enable_tiling: + pipe.vae.enable_tiling() + + # Load LoRA weights + lora_scaling = args.lora_alpha / args.rank + pipe.load_lora_weights(args.output_dir, adapter_name="cogvideox-lora") + pipe.set_adapters(["cogvideox-lora"], [lora_scaling]) + + # Run inference + validation_outputs = [] + if args.validation_prompt and args.num_validation_videos > 0: + validation_prompts = args.validation_prompt.split(args.validation_prompt_separator) + for validation_prompt in validation_prompts: + pipeline_args = { + "prompt": validation_prompt, + "guidance_scale": args.guidance_scale, + "use_dynamic_cfg": args.use_dynamic_cfg, + "height": args.height, + "width": args.width, + } + + video = log_validation( + accelerator=accelerator, + pipe=pipe, + args=args, + pipeline_args=pipeline_args, + epoch=epoch, + is_final_validation=True, + ) + validation_outputs.extend(video) + + accelerator.print("===== Memory after testing =====") + print_memory(accelerator.device) + reset_memory(accelerator.device) + torch.cuda.synchronize(accelerator.device) + + if args.push_to_hub: + save_model_card( + repo_id, + videos=validation_outputs, + base_model=args.pretrained_model_name_or_path, + validation_prompt=args.validation_prompt, + repo_folder=args.output_dir, + fps=args.fps, + ) + upload_folder( + repo_id=repo_id, + folder_path=args.output_dir, + commit_message="End of training", + ignore_patterns=["step_*", "epoch_*"], + ) + + accelerator.end_training() + + +if __name__ == "__main__": + args = get_args() + main(args) diff --git a/training/cogvideox_text_to_video_sft.py b/training/cogvideox_text_to_video_sft.py new file mode 100644 index 0000000..c2a9d8b --- /dev/null +++ b/training/cogvideox_text_to_video_sft.py @@ -0,0 +1,833 @@ +# Copyright 2024 The HuggingFace Team. +# All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import gc +import logging +import math +import os +import shutil +from datetime import timedelta +from pathlib import Path +from typing import Any, Dict + +import diffusers +import torch +import transformers +import wandb +from accelerate import Accelerator +from accelerate.logging import get_logger +from accelerate.utils import ( + DistributedDataParallelKwargs, + InitProcessGroupKwargs, + ProjectConfiguration, + set_seed, +) +from diffusers import ( + AutoencoderKLCogVideoX, + CogVideoXDPMScheduler, + CogVideoXPipeline, + CogVideoXTransformer3DModel, +) +from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution +from diffusers.optimization import get_scheduler +from diffusers.training_utils import cast_training_params +from diffusers.utils import export_to_video, is_wandb_available +from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card +from diffusers.utils.torch_utils import is_compiled_module +from huggingface_hub import create_repo, upload_folder +from torch.utils.data import DataLoader +from tqdm.auto import tqdm +from transformers import AutoTokenizer, T5EncoderModel + + +from args import get_args # isort:skip +from dataset import BucketSampler, VideoDatasetWithResizing # isort:skip +from text_encoder import compute_prompt_embeddings # isort:skip +from utils import get_gradient_norm, get_optimizer, prepare_rotary_positional_embeddings, print_memory, reset_memory # isort:skip + + +logger = get_logger(__name__) + + +def save_model_card( + repo_id: str, + videos=None, + base_model: str = None, + validation_prompt=None, + repo_folder=None, + fps=8, +): + widget_dict = [] + if videos is not None: + for i, video in enumerate(videos): + export_to_video(video, os.path.join(repo_folder, f"final_video_{i}.mp4", fps=fps)) + widget_dict.append( + { + "text": validation_prompt if validation_prompt else " ", + "output": {"url": f"video_{i}.mp4"}, + } + ) + + model_description = """TODO""" + model_card = load_or_create_model_card( + repo_id_or_path=repo_id, + from_training=True, + license="other", + base_model=base_model, + prompt=validation_prompt, + model_description=model_description, + widget=widget_dict, + ) + tags = [ + "text-to-video", + "diffusers-training", + "diffusers", + "cogvideox", + "cogvideox-diffusers", + ] + + model_card = populate_model_card(model_card, tags=tags) + model_card.save(os.path.join(repo_folder, "README.md")) + + +def log_validation( + accelerator: Accelerator, + pipe: CogVideoXPipeline, + args: Dict[str, Any], + pipeline_args: Dict[str, Any], + epoch, + is_final_validation: bool = False, +): + logger.info( + f"Running validation... \n Generating {args.num_validation_videos} videos with prompt: {pipeline_args['prompt']}." + ) + + pipe = pipe.to(accelerator.device) + + # run inference + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) if args.seed else None + + videos = [] + for _ in range(args.num_validation_videos): + video = pipe(**pipeline_args, generator=generator, output_type="np").frames[0] + videos.append(video) + + for tracker in accelerator.trackers: + phase_name = "test" if is_final_validation else "validation" + if tracker.name == "wandb": + video_filenames = [] + for i, video in enumerate(videos): + prompt = ( + pipeline_args["prompt"][:25] + .replace(" ", "_") + .replace(" ", "_") + .replace("'", "_") + .replace('"', "_") + .replace("/", "_") + ) + filename = os.path.join(args.output_dir, f"{phase_name}_video_{i}_{prompt}.mp4") + export_to_video(video, filename, fps=8) + video_filenames.append(filename) + + tracker.log( + { + phase_name: [ + wandb.Video(filename, caption=f"{i}: {pipeline_args['prompt']}") + for i, filename in enumerate(video_filenames) + ] + } + ) + + return videos + + +def main(args): + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `huggingface-cli login` to authenticate with the Hub." + ) + + if torch.backends.mps.is_available() and args.mixed_precision == "bf16": + # due to pytorch#99272, MPS does not yet support bfloat16. + raise ValueError( + "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." + ) + + logging_dir = Path(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + init_process_group_kwargs = InitProcessGroupKwargs(backend="nccl", timeout=timedelta(seconds=args.nccl_timeout)) + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + kwargs_handlers=[ddp_kwargs, init_process_group_kwargs], + ) + + # Disable AMP for MPS. + if torch.backends.mps.is_available(): + accelerator.native_amp = False + + if args.report_to == "wandb": + if not is_wandb_available(): + raise ImportError("Make sure to install wandb if you want to use it for logging during training.") + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + if args.push_to_hub: + repo_id = create_repo( + repo_id=args.hub_model_id or Path(args.output_dir).name, + exist_ok=True, + ).repo_id + + # Prepare models and scheduler + tokenizer = AutoTokenizer.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="tokenizer", + revision=args.revision, + ) + + text_encoder = T5EncoderModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="text_encoder", + revision=args.revision, + ) + + # CogVideoX-2b weights are stored in float16 + # CogVideoX-5b and CogVideoX-5b-I2V weights are stored in bfloat16 + load_dtype = torch.bfloat16 if "5b" in args.pretrained_model_name_or_path.lower() else torch.float16 + transformer = CogVideoXTransformer3DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + torch_dtype=load_dtype, + revision=args.revision, + variant=args.variant, + ) + + vae = AutoencoderKLCogVideoX.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae", + revision=args.revision, + variant=args.variant, + ) + + scheduler = CogVideoXDPMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + + if args.enable_slicing: + vae.enable_slicing() + if args.enable_tiling: + vae.enable_tiling() + + text_encoder.requires_grad_(False) + vae.requires_grad_(False) + transformer.requires_grad_(True) + + VAE_SCALING_FACTOR = vae.config.scaling_factor + VAE_SCALE_FACTOR_SPATIAL = 2 ** (len(vae.config.block_out_channels) - 1) + + # For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.state.deepspeed_plugin: + # DeepSpeed is handling precision, use what's in the DeepSpeed config + if ( + "fp16" in accelerator.state.deepspeed_plugin.deepspeed_config + and accelerator.state.deepspeed_plugin.deepspeed_config["fp16"]["enabled"] + ): + weight_dtype = torch.float16 + if ( + "bf16" in accelerator.state.deepspeed_plugin.deepspeed_config + and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"] + ): + weight_dtype = torch.float16 + else: + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + + if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: + # due to pytorch#99272, MPS does not yet support bfloat16. + raise ValueError( + "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." + ) + + text_encoder.to(accelerator.device, dtype=weight_dtype) + transformer.to(accelerator.device, dtype=weight_dtype) + vae.to(accelerator.device, dtype=weight_dtype) + + if args.gradient_checkpointing: + transformer.enable_gradient_checkpointing() + + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + for model in models: + if isinstance(model, type(unwrap_model(transformer))): + model: CogVideoXTransformer3DModel + model.save_pretrained( + os.path.join(output_dir, "transformer"), safe_serialization=True, max_shard_size="5GB" + ) + else: + raise ValueError(f"Unexpected save model: {model.__class__}") + + # make sure to pop weight so that corresponding model is not saved again + weights.pop() + + def load_model_hook(models, input_dir): + transformer_ = None + + while len(models) > 0: + model = models.pop() + + if isinstance(model, type(unwrap_model(transformer))): + transformer_: CogVideoXTransformer3DModel = model + else: + raise ValueError(f"Unexpected save model: {model.__class__.__name__}") + + load_model = CogVideoXTransformer3DModel.from_pretrained(os.path.join(input_dir, "transformer")) + transformer_.register_to_config(**load_model.config) + transformer_.load_state_dict(load_model.state_dict()) + del load_model + + # Make sure the trainable params are in float32. This is again needed since the base models + # are in `weight_dtype`. More details: + # https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804 + if args.mixed_precision == "fp16": + cast_training_params([transformer_]) + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32 and torch.cuda.is_available(): + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Make sure the trainable params are in float32. + if args.mixed_precision == "fp16": + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params([transformer], dtype=torch.float32) + + transformer_parameters = list(filter(lambda p: p.requires_grad, transformer.parameters())) + + # Optimization parameters + transformer_parameters_with_lr = { + "params": transformer_parameters, + "lr": args.learning_rate, + } + params_to_optimize = [transformer_parameters_with_lr] + + use_deepspeed_optimizer = ( + accelerator.state.deepspeed_plugin is not None + and "optimizer" in accelerator.state.deepspeed_plugin.deepspeed_config + ) + use_deepspeed_scheduler = ( + accelerator.state.deepspeed_plugin is not None + and "scheduler" not in accelerator.state.deepspeed_plugin.deepspeed_config + ) + + optimizer = get_optimizer( + params_to_optimize=params_to_optimize, + optimizer_name=args.optimizer, + learning_rate=args.learning_rate, + beta1=args.beta1, + beta2=args.beta2, + beta3=args.beta3, + epsilon=args.epsilon, + weight_decay=args.weight_decay, + prodigy_decouple=args.prodigy_decouple, + prodigy_use_bias_correction=args.prodigy_use_bias_correction, + prodigy_safeguard_warmup=args.prodigy_safeguard_warmup, + use_8bit=args.use_8bit, + use_deepspeed=use_deepspeed_optimizer, + ) + + # Dataset and DataLoader + train_dataset = VideoDatasetWithResizing( + data_root=args.data_root, + dataset_file=args.dataset_file, + caption_column=args.caption_column, + video_column=args.video_column, + max_num_frames=args.max_num_frames, + id_token=args.id_token, + height_buckets=args.height_buckets, + width_buckets=args.width_buckets, + frame_buckets=args.frame_buckets, + load_tensors=args.load_tensors, + random_flip=args.random_flip, + ) + + def collate_fn_without_pre_encoding(data): + prompts = [x["prompt"] for x in data[0]] + + videos = [x["video"] for x in data[0]] + videos = torch.stack(videos) + videos = videos.to(accelerator.device, dtype=weight_dtype) + videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W] + latent_dist = vae.encode(videos).latent_dist + videos = latent_dist.sample() * VAE_SCALING_FACTOR + videos = videos.permute(0, 2, 1, 3, 4) # [B, F, C, H, W] + videos = videos.to(memory_format=torch.contiguous_format).float() + + return { + "videos": videos, + "prompts": prompts, + } + + def collate_fn_with_pre_encoding(data): + prompts = [x["prompt"] for x in data[0]] + prompts = torch.stack(prompts).to(accelerator.device, dtype=weight_dtype) + + videos = [x["video"] for x in data[0]] + videos = torch.stack(videos).to(accelerator.device, dtype=weight_dtype) + videos = DiagonalGaussianDistribution(videos).sample() * VAE_SCALING_FACTOR + videos = videos.permute(0, 2, 1, 3, 4) # [B, F, C, H, W] + videos = videos.to(memory_format=torch.contiguous_format).float() + + return { + "videos": videos, + "prompts": prompts, + } + + train_dataloader = DataLoader( + train_dataset, + batch_size=1, + sampler=BucketSampler(train_dataset, batch_size=args.train_batch_size, shuffle=True), + collate_fn=collate_fn_with_pre_encoding if args.load_tensors else collate_fn_without_pre_encoding, + num_workers=args.dataloader_num_workers, + ) + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + if use_deepspeed_scheduler: + from accelerate.utils import DummyScheduler + + lr_scheduler = DummyScheduler( + name=args.lr_scheduler, + optimizer=optimizer, + total_num_steps=args.max_train_steps * accelerator.num_processes, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + ) + else: + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + num_cycles=args.lr_num_cycles, + power=args.lr_power, + ) + + # Prepare everything with our `accelerator`. + transformer, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer, optimizer, train_dataloader, lr_scheduler + ) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataset) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_name = args.tracker_name or "cogvideox-sft" + accelerator.init_trackers(tracker_name, config=vars(args)) + + accelerator.print("===== Memory before training =====") + reset_memory(accelerator.device) + print_memory(accelerator.device) + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + num_trainable_parameters = sum(param.numel() for model in params_to_optimize for param in model["params"]) + + accelerator.print("***** Running training *****") + accelerator.print(f" Num trainable parameters = {num_trainable_parameters}") + accelerator.print(f" Num examples = {len(train_dataset)}") + accelerator.print(f" Num epochs = {args.num_train_epochs}") + accelerator.print(f" Instantaneous batch size per device = {args.train_batch_size}") + accelerator.print(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + accelerator.print(f" Gradient accumulation steps = {args.gradient_accumulation_steps}") + accelerator.print(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if not args.resume_from_checkpoint: + initial_global_step = 0 + else: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the mos recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + first_epoch = global_step // num_update_steps_per_epoch + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + # For DeepSpeed training + model_config = transformer.module.config if hasattr(transformer, "module") else transformer.config + + if args.load_tensors: + del vae, text_encoder + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize(accelerator.device) + + for epoch in range(first_epoch, args.num_train_epochs): + transformer.train() + + for step, batch in enumerate(train_dataloader): + models_to_accumulate = [transformer] + + with accelerator.accumulate(models_to_accumulate): + model_input = batch["videos"] + prompts = batch["prompts"] + + # Encode prompts + if not args.load_tensors: + prompt_embeds = compute_prompt_embeddings( + tokenizer, + text_encoder, + prompts, + model_config.max_text_seq_length, + accelerator.device, + weight_dtype, + requires_grad=False, + ) + else: + prompt_embeds = prompts + + # Sample noise that will be added to the latents + noise = torch.randn_like(model_input) + batch_size, num_frames, num_channels, height, width = model_input.shape + + # Sample a random timestep for each image + timesteps = torch.randint( + 0, + scheduler.config.num_train_timesteps, + (batch_size,), + dtype=torch.int64, + device=model_input.device, + ) + + # Prepare rotary embeds + image_rotary_emb = ( + prepare_rotary_positional_embeddings( + height=height * VAE_SCALE_FACTOR_SPATIAL, + width=width * VAE_SCALE_FACTOR_SPATIAL, + num_frames=num_frames, + vae_scale_factor_spatial=VAE_SCALE_FACTOR_SPATIAL, + patch_size=model_config.patch_size, + attention_head_dim=model_config.attention_head_dim, + device=accelerator.device, + ) + if model_config.use_rotary_positional_embeddings + else None + ) + + # Add noise to the model input according to the noise magnitude at each timestep + # (this is the forward diffusion process) + noisy_model_input = scheduler.add_noise(model_input, noise, timesteps) + + # Predict the noise residual + model_output = transformer( + hidden_states=noisy_model_input, + encoder_hidden_states=prompt_embeds, + timestep=timesteps, + image_rotary_emb=image_rotary_emb, + return_dict=False, + )[0] + + model_pred = scheduler.get_velocity(model_output, noisy_model_input, timesteps) + + alphas_cumprod = scheduler.alphas_cumprod[timesteps] + weights = 1 / (1 - alphas_cumprod) + while len(weights.shape) < len(model_pred.shape): + weights = weights.unsqueeze(-1) + + target = model_input + + loss = torch.mean( + (weights * (model_pred - target) ** 2).reshape(batch_size, -1), + dim=1, + ) + loss = loss.mean() + accelerator.backward(loss) + + if accelerator.sync_gradients: + gradient_norm_before_clip = get_gradient_norm(transformer.parameters()) + accelerator.clip_grad_norm_(transformer.parameters(), args.max_grad_norm) + gradient_norm_after_clip = get_gradient_norm(transformer.parameters()) + + if accelerator.state.deepspeed_plugin is None: + optimizer.step() + optimizer.zero_grad() + + lr_scheduler.step() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + global_step += 1 + + if accelerator.is_main_process: + if global_step % args.checkpointing_steps == 0: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"Removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + logs = { + "loss": loss.detach().item(), + "lr": lr_scheduler.get_last_lr()[0], + "gradient_norm_before_clip": gradient_norm_before_clip, + "gradient_norm_after_clip": gradient_norm_after_clip, + } + progress_bar.set_postfix(**logs) + accelerator.log(logs, step=global_step) + + if global_step >= args.max_train_steps: + break + + if accelerator.is_main_process: + if args.validation_prompt is not None and (epoch + 1) % args.validation_epochs == 0: + accelerator.print("===== Memory before validation =====") + print_memory(accelerator.device) + torch.cuda.synchronize(accelerator.device) + + pipe = CogVideoXPipeline.from_pretrained( + args.pretrained_model_name_or_path, + transformer=unwrap_model(transformer), + scheduler=scheduler, + revision=args.revision, + variant=args.variant, + torch_dtype=weight_dtype, + ) + + if args.enable_slicing: + pipe.vae.enable_slicing() + if args.enable_tiling: + pipe.vae.enable_tiling() + + validation_prompts = args.validation_prompt.split(args.validation_prompt_separator) + for validation_prompt in validation_prompts: + pipeline_args = { + "prompt": validation_prompt, + "guidance_scale": args.guidance_scale, + "use_dynamic_cfg": args.use_dynamic_cfg, + "height": args.height, + "width": args.width, + } + + log_validation( + accelerator=accelerator, + pipe=pipe, + args=args, + pipeline_args=pipeline_args, + epoch=epoch, + is_final_validation=False, + ) + + accelerator.print("===== Memory after validation =====") + print_memory(accelerator.device) + reset_memory(accelerator.device) + + del pipe + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize(accelerator.device) + + accelerator.wait_for_everyone() + + if accelerator.is_main_process: + transformer = unwrap_model(transformer) + dtype = ( + torch.float16 + if args.mixed_precision == "fp16" + else torch.bfloat16 + if args.mixed_precision == "bf16" + else torch.float32 + ) + transformer = transformer.to(dtype) + + transformer.save_pretrained( + os.path.join(args.output_dir, "transformer"), + safe_serialization=True, + max_shard_size="5GB", + ) + + # Cleanup trained models to save memory + if args.load_tensors: + del transformer + else: + del transformer, text_encoder, vae + + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize(accelerator.device) + + accelerator.print("===== Memory before testing =====") + print_memory(accelerator.device) + reset_memory(accelerator.device) + + # Final test inference + pipe = CogVideoXPipeline.from_pretrained( + args.pretrained_model_name_or_path, + revision=args.revision, + variant=args.variant, + torch_dtype=weight_dtype, + ) + pipe.scheduler = CogVideoXDPMScheduler.from_config(pipe.scheduler.config) + + if args.enable_slicing: + pipe.vae.enable_slicing() + if args.enable_tiling: + pipe.vae.enable_tiling() + + # Run inference + validation_outputs = [] + if args.validation_prompt and args.num_validation_videos > 0: + validation_prompts = args.validation_prompt.split(args.validation_prompt_separator) + for validation_prompt in validation_prompts: + pipeline_args = { + "prompt": validation_prompt, + "guidance_scale": args.guidance_scale, + "use_dynamic_cfg": args.use_dynamic_cfg, + "height": args.height, + "width": args.width, + } + + video = log_validation( + accelerator=accelerator, + pipe=pipe, + args=args, + pipeline_args=pipeline_args, + epoch=epoch, + is_final_validation=True, + ) + validation_outputs.extend(video) + + accelerator.print("===== Memory after testing =====") + print_memory(accelerator.device) + reset_memory(accelerator.device) + torch.cuda.synchronize(accelerator.device) + + if args.push_to_hub: + save_model_card( + repo_id, + videos=validation_outputs, + base_model=args.pretrained_model_name_or_path, + validation_prompt=args.validation_prompt, + repo_folder=args.output_dir, + fps=args.fps, + ) + upload_folder( + repo_id=repo_id, + folder_path=args.output_dir, + commit_message="End of training", + ignore_patterns=["step_*", "epoch_*"], + ) + + accelerator.end_training() + + +if __name__ == "__main__": + args = get_args() + main(args) diff --git a/training/dataset.py b/training/dataset.py new file mode 100644 index 0000000..fe125c9 --- /dev/null +++ b/training/dataset.py @@ -0,0 +1,289 @@ +import random +from pathlib import Path +from typing import Any, Dict, List, Optional, Tuple + +import pandas as pd +import torch +from accelerate.logging import get_logger +from torch.utils.data import Dataset, Sampler +from torchvision import transforms +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__) + +HEIGHT_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536] +WIDTH_BUCKETS = [256, 320, 384, 480, 512, 576, 720, 768, 960, 1024, 1280, 1536] +FRAME_BUCKETS = [16, 24, 32, 48, 64, 80] + + +class VideoDataset(Dataset): + def __init__( + self, + data_root: str, + dataset_file: Optional[str] = None, + caption_column: str = "text", + video_column: str = "video", + max_num_frames: int = 49, + id_token: Optional[str] = None, + height_buckets: List[int] = None, + width_buckets: List[int] = None, + frame_buckets: List[int] = None, + load_tensors: bool = False, + random_flip: Optional[float] = None, + ) -> None: + super().__init__() + + self.data_root = Path(data_root) + self.dataset_file = dataset_file + self.caption_column = caption_column + self.video_column = video_column + self.max_num_frames = max_num_frames + self.id_token = id_token or "" + self.height_buckets = height_buckets or HEIGHT_BUCKETS + self.width_buckets = width_buckets or WIDTH_BUCKETS + self.frame_buckets = frame_buckets or FRAME_BUCKETS + self.load_tensors = load_tensors + self.random_flip = random_flip + + self.resolutions = [ + (f, h, w) for h in self.height_buckets for w in self.width_buckets for f in self.frame_buckets + ] + + # Two methods of loading data are supported. + # - Using a CSV: caption_column and video_column must be some column in the CSV. One could + # make use of other columns too, such as a motion score or aesthetic score, by modifying the + # logic in CSV processing. + # - Using two files containing line-separate captions and relative paths to videos. + # For a more detailed explanation about preparing dataset format, checkout the README. + if dataset_file is None: + ( + self.prompts, + self.video_paths, + ) = self._load_dataset_from_local_path() + else: + ( + self.prompts, + self.video_paths, + ) = self._load_dataset_from_csv() + + self.num_videos = len(self.video_paths) + if self.num_videos != len(self.prompts): + raise ValueError( + f"Expected length of prompts and videos to be the same but found {len(self.prompts)=} and {len(self.video_paths)=}. Please ensure that the number of caption prompts and videos match in your dataset." + ) + + self.video_transforms = transforms.Compose( + [ + transforms.RandomHorizontalFlip(random_flip) if random_flip else transforms.Lambda(lambda x: x), + transforms.Lambda(lambda x: x / 255.0), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ] + ) + + def __len__(self) -> int: + return self.num_videos + + def __getitem__(self, index: int) -> Dict[str, Any]: + if isinstance(index, list): + # Here, index is actually a list of data objects that we need to return. + # The BucketSampler should ideally return indices. But, in the sampler, we'd like + # to have information about num_frames, height and width. Since this is not stored + # as metadata, we need to read the video to get this information. You could read this + # information without loading the full video in memory, but we do it anyway. In order + # to not load the video twice (once to get the metadata, and once to return the loaded video + # based on sampled indices), we cache it in the BucketSampler. When the sampler is + # to yield, we yield the cache data instead of indices. So, this special check ensures + # that data is not loaded a second time. PRs are welcome for improvements. + return index + + if self.load_tensors: + latents, prompt_embeds = self._preprocess_video(self.video_paths[index]) + + # This is hardcoded for now. + # The VAE's temporal compression ratio is 4. + # The VAE's spatial compression ratio is 8. + latent_num_frames = latents.size(1) + if latent_num_frames % 2 == 0: + num_frames = latent_num_frames * 4 + else: + num_frames = (latent_num_frames - 1) * 4 + 1 + + height = latents.size(2) * 8 + width = latents.size(3) * 8 + + return { + "prompt": prompt_embeds, + "video": latents, + "video_metadata": { + "num_frames": num_frames, + "height": height, + "width": width, + }, + } + else: + video, _ = self._preprocess_video(self.video_paths[index]) + + return { + "prompt": self.id_token + self.prompts[index], + "video": video, + "video_metadata": { + "num_frames": video.shape[0], + "height": video.shape[2], + "width": video.shape[3], + }, + } + + def _load_dataset_from_local_path(self) -> Tuple[List[str], List[str]]: + if not self.data_root.exists(): + raise ValueError("Root folder for videos does not exist") + + prompt_path = self.data_root.joinpath(self.caption_column) + video_path = self.data_root.joinpath(self.video_column) + + if not prompt_path.exists() or not prompt_path.is_file(): + raise ValueError( + "Expected `--caption_column` to be path to a file in `--data_root` containing line-separated text prompts." + ) + if not video_path.exists() or not video_path.is_file(): + raise ValueError( + "Expected `--video_column` to be path to a file in `--data_root` containing line-separated paths to video data in the same directory." + ) + + with open(prompt_path, "r", encoding="utf-8") as file: + prompts = [line.strip() for line in file.readlines() if len(line.strip()) > 0] + with open(video_path, "r", encoding="utf-8") as file: + video_paths = [self.data_root.joinpath(line.strip()) for line in file.readlines() if len(line.strip()) > 0] + + if not self.load_tensors and any(not path.is_file() for path in video_paths): + raise ValueError( + f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file." + ) + + return prompts, video_paths + + def _load_dataset_from_csv(self) -> Tuple[List[str], List[str]]: + df = pd.read_csv(self.dataset_file) + prompts = df[self.caption_column].tolist() + video_paths = df[self.video_column].tolist() + video_paths = [self.data_root.joinpath(line.strip()) for line in video_paths] + + if any(not path.is_file() for path in video_paths): + raise ValueError( + f"Expected `{self.video_column=}` to be a path to a file in `{self.data_root=}` containing line-separated paths to video data but found atleast one path that is not a valid file." + ) + + return prompts, video_paths + + def _preprocess_video(self, path: Path) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + r""" + Loads a single video, or latent and prompt embedding, based on initialization parameters. + + If returning a video, returns a [F, C, H, W] video tensor, and None for the prompt embedding. Here, + F, C, H and W are the frames, channels, height and width of the input video. + + If returning latent/embedding, returns a [F, C, H, W] latent, and the prompt embedding of shape [S, D]. + F, C, H and W are the frames, channels, height and width of the latent, and S, D are the sequence length + and embedding dimension of prompt embeddings. + """ + if self.load_tensors: + return self._load_preprocessed_latents_and_embeds(path) + else: + video_reader = decord.VideoReader(uri=path.as_posix()) + video_num_frames = len(video_reader) + + indices = list(range(0, video_num_frames, video_num_frames // self.max_num_frames)) + frames = video_reader.get_batch(indices) + frames = frames[: self.max_num_frames].float() + frames = frames.permute(0, 3, 1, 2).contiguous() + frames = torch.stack([self.video_transforms(frame) for frame in frames], dim=0) + + return frames, None + + def _load_preprocessed_latents_and_embeds(self, path: Path) -> Tuple[torch.Tensor, torch.Tensor]: + filename_without_ext = path.name.split(".")[0] + pt_filename = f"{filename_without_ext}.pt" + + # The current path is something like: /a/b/c/d/videos/00001.mp4 + # We need to reach: /a/b/c/d/latents/00001.pt + latents_path = path.parent.parent.joinpath("latents") + embeds_path = path.parent.parent.joinpath("embeddings") + + if not latents_path.exists() or not embeds_path.exists(): + raise ValueError( + f"When setting the load_tensors parameter to `True`, it is expected that the `{self.data_root=}` contains two folders named `latents` and `embeddings`. However, these folders were not found. Please make sure to have prepared your data correctly using `prepare_data.py`." + ) + + latent_filepath = latents_path.joinpath(pt_filename) + embeds_filepath = embeds_path.joinpath(pt_filename) + + if not latent_filepath.is_file() or not embeds_filepath.is_file(): + latent_filepath = latent_filepath.as_posix() + embeds_filepath = embeds_filepath.as_posix() + raise ValueError( + f"The file {latent_filepath=} or {embeds_filepath=} could not be found. Please ensure that you've correctly executed `prepare_dataset.py`." + ) + + latents = torch.load(latent_filepath, map_location="cpu", weights_only=True) + embeds = torch.load(embeds_filepath, map_location="cpu", weights_only=True) + + return latents, embeds + + +class VideoDatasetWithResizing(VideoDataset): + def __init__(self, *args, **kwargs) -> None: + super().__init__(*args, **kwargs) + + def _preprocess_video(self, path: Path) -> torch.Tensor: + if self.load_tensors: + return self._load_preprocessed_latents_and_embeds(path) + else: + video_reader = decord.VideoReader(uri=path.as_posix()) + video_num_frames = len(video_reader) + nearest_frame_bucket = min( + self.frame_buckets, key=lambda x: abs(x - min(video_num_frames, self.max_num_frames)) + ) + + 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, None + + def _find_nearest_resolution(self, height, width): + nearest_res = min(self.resolutions, key=lambda x: abs(x[1] - height) + abs(x[2] - width)) + return nearest_res[1], nearest_res[2] + + +class BucketSampler(Sampler): + def __init__(self, data_source: VideoDataset, batch_size: int = 8, shuffle: bool = True) -> None: + self.data_source = data_source + self.batch_size = batch_size + self.shuffle = shuffle + + self.buckets = {resolution: [] for resolution in data_source.resolutions} + + 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)] = [] diff --git a/training/prepare_dataset.py b/training/prepare_dataset.py new file mode 100644 index 0000000..9d70884 --- /dev/null +++ b/training/prepare_dataset.py @@ -0,0 +1,458 @@ +#!/usr/bin/env python3 + +# For folder structure dataset: python3 prepare_dataset.py --model_id THUDM/CogVideoX-2b --data_root /raid/aryan/video-dataset-disney/ --caption_column prompts.txt --video_column videos.txt --output_dir dump --height 480 --width 720 --max_num_frames 49 --max_sequence_length 226 --target_fps 8 --batch_size 1 --dtype fp32 +# For latent/embed structure dataset: python3 prepare_dataset.py --model_id THUDM/CogVideoX-2b --data_root /raid/aryan/video-dataset-disney/ --caption_column prompts.txt --video_column videos.txt --output_dir dump --height 480 --width 720 --max_num_frames 49 --max_sequence_length 226 --target_fps 8 --batch_size 1 --dtype fp32 --save_tensors + +import argparse +import gc +import pathlib +import traceback +from typing import Any, Dict, List, Optional, Tuple, Union + +import pandas as pd +import torch +from diffusers import AutoencoderKLCogVideoX +from diffusers.utils import export_to_video, get_logger +from torchvision import transforms +from transformers import T5EncoderModel, T5Tokenizer + + +# Must import after importing torch, otherwise there's a nasty segfault when loading text_encoder/vae +import decord # isort:skip + +decord.bridge.set_bridge("torch") + +logger = get_logger(__name__) + +DTYPE_MAPPING = { + "fp32": torch.float32, + "fp16": torch.float16, + "bf16": torch.bfloat16, +} + + +def get_args() -> Dict[str, Any]: + parser = argparse.ArgumentParser() + parser.add_argument( + "--model_id", + type=str, + default="THUDM/CogVideoX-2b", + help="Hugging Face model ID to use for tokenizer, text encoder and VAE.", + ) + parser.add_argument("--data_root", type=str, required=True, help="Path to where training data is located.") + parser.add_argument( + "--dataset_file", type=str, default=None, help="Path to CSV file containing metadata about training data." + ) + parser.add_argument( + "--caption_column", + type=str, + default="caption", + help="If using a CSV file via the `--dataset_file` argument, this should be the name of the column containing the captions. If using the folder structure format for data loading, this should be the name of the file containing line-separated captions (the file should be located in `--data_root`).", + ) + parser.add_argument( + "--video_column", + type=str, + default="video", + help="If using a CSV file via the `--dataset_file` argument, this should be the name of the column containing the video paths. If using the folder structure format for data loading, this should be the name of the file containing line-separated video paths (the file should be located in `--data_root`).", + ) + parser.add_argument( + "--output_dir", + type=str, + required=True, + help="Path to output directory where preprocessed videos/latents/embeddings will be saved.", + ) + parser.add_argument("--height", type=int, default=480, help="Height of the resized output video.") + parser.add_argument("--width", type=int, default=720, help="Width of the resized output video.") + parser.add_argument("--max_num_frames", type=int, default=49, help="Maximum number of frames in output video.") + parser.add_argument( + "--max_sequence_length", type=int, default=226, help="Max sequence length of prompt embeddings." + ) + parser.add_argument( + "--target_fps", type=int, default=8, help="Frame rate of output videos if `--save_tensors` is unspecified." + ) + parser.add_argument( + "--save_tensors", + action="store_true", + help="Whether to encode videos/captions to latents/embeddings and save them in pytorch serializable format.", + ) + parser.add_argument( + "--use_slicing", + action="store_true", + help="Whether to enable sliced encoding/decoding in the VAE. Only used if `--save_tensors` is also used.", + ) + parser.add_argument( + "--use_tiling", + action="store_true", + help="Whether to enable tiled encoding/decoding in the VAE. Only used if `--save_tensors` is also used.", + ) + parser.add_argument("--batch_size", type=int, default=1, help="Number of videos to process at once in the VAE.") + parser.add_argument( + "--num_decode_threads", + type=int, + default=0, + help="Number of decoding threads for `decord` to use. The default `0` means to automatically determine required number of threads.", + ) + parser.add_argument( + "--dtype", + type=str, + choices=["fp32", "fp16", "bf16"], + default="fp32", + help="Data type to use when generating latents and prompt embeddings.", + ) + return parser.parse_args() + + +def load_dataset_from_local_path( + data_root: pathlib.Path, caption_column: str, video_column: str +) -> Tuple[List[str], List[str]]: + if not data_root.exists(): + raise ValueError("Root folder for videos does not exist") + + prompt_path = data_root.joinpath(caption_column) + video_path = data_root.joinpath(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 = [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 `{video_column=}` to be a path to a file in `{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( + data_root: pathlib.Path, dataset_file: pathlib.Path, caption_column: str, video_column: str +) -> Tuple[List[str], List[str]]: + df = pd.read_csv(dataset_file) + prompts = df[caption_column].tolist() + video_paths = df[video_column].tolist() + video_paths = [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 `{video_column=}` to be a path to a file in `{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_and_preprocess_video( + path: pathlib.Path, height: int, width: int, max_num_frames: int, video_transforms, num_threads: int = 0 +) -> torch.Tensor: + frames = None + + try: + video_reader = decord.VideoReader(uri=path.as_posix(), height=height, width=width, num_threads=num_threads) + video_num_frames = len(video_reader) + + if video_num_frames < max_num_frames: + logger.warning( + f"Video at `{path.as_posix()}` should have atleast `{max_num_frames=}`, but got only `{video_num_frames=}`. Skipping it." + ) + return + + indices = list(range(0, video_num_frames, video_num_frames // max_num_frames)) + frames: torch.Tensor = video_reader.get_batch(indices) + frames = frames[:max_num_frames].float() + frames = frames.permute(0, 3, 1, 2).contiguous() + frames = torch.stack([video_transforms(frame) for frame in frames], dim=0) + except Exception as e: + logger.error(f"Error: {e}. Skipping video located at `{path.as_posix()}`") + traceback.print_exc() + + return frames + + +def _get_t5_prompt_embeds( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: Union[str, List[str]], + num_videos_per_prompt: int = 1, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + text_input_ids=None, +): + prompt = [prompt] if isinstance(prompt, str) else prompt + batch_size = len(prompt) + + if tokenizer is not None: + 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 + else: + if text_input_ids is None: + raise ValueError("`text_input_ids` must be provided when the tokenizer is not specified.") + + prompt_embeds = text_encoder(text_input_ids.to(device))[0] + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + # duplicate text embeddings for each generation per prompt, using mps friendly method + _, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1) + + return prompt_embeds + + +def encode_prompt( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: Union[str, List[str]], + num_videos_per_prompt: int = 1, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + text_input_ids=None, +): + prompt = [prompt] if isinstance(prompt, str) else prompt + prompt_embeds = _get_t5_prompt_embeds( + tokenizer, + text_encoder, + prompt=prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + text_input_ids=text_input_ids, + ) + return prompt_embeds + + +def compute_prompt_embeddings( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: str, + max_sequence_length: int, + device: torch.device, + dtype: torch.dtype, + requires_grad: bool = False, +): + if requires_grad: + prompt_embeds = encode_prompt( + tokenizer, + text_encoder, + prompt, + num_videos_per_prompt=1, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + else: + with torch.no_grad(): + prompt_embeds = encode_prompt( + tokenizer, + text_encoder, + prompt, + num_videos_per_prompt=1, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + return prompt_embeds + + +def save_videos( + videos: torch.Tensor, video_paths: List[str], prompts: List[str], output_dir: pathlib.Path, target_fps: int = 8 +) -> None: + assert videos.size(0) == len(video_paths) + + videos = (videos + 1) / 2 + videos = (videos * 255.0).clip(0, 255) + videos = videos.to(dtype=torch.uint8) + + video_dir = output_dir.joinpath("videos") + + output_dir.mkdir(parents=True, exist_ok=True) + video_dir.mkdir(parents=True, exist_ok=True) + + to_pil_image = transforms.ToPILImage() + videos_pil = [[to_pil_image(frame) for frame in video] for video in videos] + + for video, video_path in zip(videos_pil, video_paths): + filename = video_dir.joinpath(pathlib.Path(video_path).name) + logger.debug(f"Saving video to `{filename}`") + export_to_video(video, filename.as_posix(), fps=target_fps) + + with open(output_dir.joinpath("videos.txt").as_posix(), "w", encoding="utf-8") as file: + for video_path in video_paths: + file.write(f"videos/{pathlib.Path(video_path).name}\n") + + with open(output_dir.joinpath("prompts.txt").as_posix(), "w", encoding="utf-8") as file: + for prompt in prompts: + file.write(f"{prompt}\n") + + +def save_latents_and_embeddings( + latents: torch.Tensor, + prompt_embeds: torch.Tensor, + video_paths: List[str], + prompts: List[str], + output_dir: pathlib.Path, +) -> None: + assert latents.size(0) == prompt_embeds.size(0) + assert latents.size(0) == len(video_paths) + assert prompt_embeds.size(0) == len(prompts) + + latents_dir = output_dir.joinpath("latents") + embeds_dir = output_dir.joinpath("embeddings") + + output_dir.mkdir(parents=True, exist_ok=True) + latents_dir.mkdir(parents=True, exist_ok=True) + embeds_dir.mkdir(parents=True, exist_ok=True) + + for latent, embed, video_path in zip(latents, prompt_embeds, video_paths): + # Need to perform the clone, otherwise the entire `latents` or `prompt_embeds` tensor is + # saved for every single video/prompt embedding. This is due to us viewing a slice of a + # large tensor when iteratively saving stuff here. + latent = latent.clone() + embed = embed.clone() + + video_path = pathlib.Path(video_path) + filename_without_ext = video_path.name.split(".")[0] + + latent_filename = latents_dir.joinpath(filename_without_ext) + embed_filename = embeds_dir.joinpath(filename_without_ext) + + latent_filename = f"{latent_filename}.pt" + embed_filename = f"{embed_filename}.pt" + + torch.save(latent, latent_filename) + torch.save(embed, embed_filename) + + with open(output_dir.joinpath("videos.txt").as_posix(), "w", encoding="utf-8") as file: + for video_path in video_paths: + file.write(f"videos/{pathlib.Path(video_path).name}\n") + + with open(output_dir.joinpath("prompts.txt").as_posix(), "w", encoding="utf-8") as file: + for prompt in prompts: + file.write(f"{prompt}\n") + + +@torch.no_grad() +def main(args: Dict[str, Any]) -> None: + data_root = pathlib.Path(args.data_root) + dataset_file = None + if args.dataset_file: + dataset_file = pathlib.Path(args.dataset_file) + + if dataset_file is None: + prompts, video_paths = load_dataset_from_local_path(data_root, args.caption_column, args.video_column) + else: + prompts, video_paths = load_dataset_from_csv(data_root, dataset_file, args.caption_column, args.video_column) + + video_transforms = transforms.Compose( + [ + transforms.Lambda(lambda x: x / 255.0), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ] + ) + + prompts_usable = [] + video_paths_usable = [] + videos = [] + for prompt, path in zip(prompts, video_paths): + video = load_and_preprocess_video( + path, args.height, args.width, args.max_num_frames, video_transforms, args.num_decode_threads + ) + if video is not None: + prompts_usable.append(prompt) + video_paths_usable.append(path) + videos.append(video) + videos = torch.stack(videos) + + if not args.save_tensors: + save_videos(videos, video_paths_usable, prompts_usable, pathlib.Path(args.output_dir), args.target_fps) + else: + dtype = DTYPE_MAPPING[args.dtype] + tokenizer = T5Tokenizer.from_pretrained(args.model_id, subfolder="tokenizer") + text_encoder = T5EncoderModel.from_pretrained(args.model_id, subfolder="text_encoder", torch_dtype=dtype) + text_encoder = text_encoder.to("cuda") + + prompt_embeds_list = [] + for start_index in range(0, len(prompts_usable), args.batch_size): + end_index = min(len(prompts_usable), start_index + args.batch_size) + batch_prompts = prompts_usable[start_index:end_index] + + prompt_embeds = compute_prompt_embeddings( + tokenizer, + text_encoder, + batch_prompts, + max_sequence_length=args.max_sequence_length, + device="cuda", + dtype=dtype, + ) + prompt_embeds_list.append(prompt_embeds) + + prompt_embeds = torch.cat(prompt_embeds_list).to("cpu") + + del tokenizer, text_encoder + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize("cuda") + + vae = AutoencoderKLCogVideoX.from_pretrained(args.model_id, subfolder="vae", torch_dtype=dtype) + vae = vae.to("cuda") + + if args.use_slicing: + vae.enable_slicing() + if args.use_tiling: + vae.enable_tiling() + + encoded_videos = [] + for start_index in range(0, len(video_paths_usable), args.batch_size): + end_index = min(len(video_paths_usable), start_index + args.batch_size) + batch_videos = videos[start_index:end_index] + + batch_videos = batch_videos.to("cuda") + batch_videos = batch_videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W] + + if args.use_slicing: + encoded_slices = [vae._encode(video_slice) for video_slice in batch_videos.split(1)] + encoded_video = torch.cat(encoded_slices) + else: + encoded_video = vae._encode(batch_videos) + + encoded_videos.append(encoded_video) + + encoded_videos = torch.cat(encoded_videos).to("cpu") + + del vae + gc.collect() + torch.cuda.empty_cache() + torch.cuda.synchronize("cuda") + + save_latents_and_embeddings( + encoded_videos, prompt_embeds, video_paths_usable, prompts_usable, pathlib.Path(args.output_dir) + ) + + +if __name__ == "__main__": + args = get_args() + + assert args.height % 16 == 0, "CogVideoX requires input video height to be divisible by 16." + assert args.width % 16 == 0, "CogVideoX requires input video width to be divisible by 16." + assert ( + args.max_num_frames % 4 == 0 or args.max_num_frames % 4 == 1 + ), "`--max_num_frames` must be of form 4 * k or 4 * k + 1 to be compatible with VAE." + + main(args) diff --git a/training/text_encoder/__init__.py b/training/text_encoder/__init__.py new file mode 100644 index 0000000..09f9e8c --- /dev/null +++ b/training/text_encoder/__init__.py @@ -0,0 +1 @@ +from .text_encoder import compute_prompt_embeddings diff --git a/training/text_encoder/text_encoder.py b/training/text_encoder/text_encoder.py new file mode 100644 index 0000000..9237875 --- /dev/null +++ b/training/text_encoder/text_encoder.py @@ -0,0 +1,99 @@ +from typing import List, Optional, Union + +import torch +from transformers import T5EncoderModel, T5Tokenizer + + +def _get_t5_prompt_embeds( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: Union[str, List[str]], + num_videos_per_prompt: int = 1, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + text_input_ids=None, +): + prompt = [prompt] if isinstance(prompt, str) else prompt + batch_size = len(prompt) + + if tokenizer is not None: + 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 + else: + if text_input_ids is None: + raise ValueError("`text_input_ids` must be provided when the tokenizer is not specified.") + + prompt_embeds = text_encoder(text_input_ids.to(device))[0] + prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) + + # duplicate text embeddings for each generation per prompt, using mps friendly method + _, seq_len, _ = prompt_embeds.shape + prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1) + prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1) + + return prompt_embeds + + +def encode_prompt( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: Union[str, List[str]], + num_videos_per_prompt: int = 1, + max_sequence_length: int = 226, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + text_input_ids=None, +): + prompt = [prompt] if isinstance(prompt, str) else prompt + prompt_embeds = _get_t5_prompt_embeds( + tokenizer, + text_encoder, + prompt=prompt, + num_videos_per_prompt=num_videos_per_prompt, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + text_input_ids=text_input_ids, + ) + return prompt_embeds + + +def compute_prompt_embeddings( + tokenizer: T5Tokenizer, + text_encoder: T5EncoderModel, + prompt: str, + max_sequence_length: int, + device: torch.device, + dtype: torch.dtype, + requires_grad: bool = False, +): + if requires_grad: + prompt_embeds = encode_prompt( + tokenizer, + text_encoder, + prompt, + num_videos_per_prompt=1, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + else: + with torch.no_grad(): + prompt_embeds = encode_prompt( + tokenizer, + text_encoder, + prompt, + num_videos_per_prompt=1, + max_sequence_length=max_sequence_length, + device=device, + dtype=dtype, + ) + return prompt_embeds diff --git a/training/utils.py b/training/utils.py new file mode 100644 index 0000000..609045f --- /dev/null +++ b/training/utils.py @@ -0,0 +1,182 @@ +import gc +from typing import Optional, Tuple, Union + +import torch +from accelerate.logging import get_logger +from diffusers.models.embeddings import get_3d_rotary_pos_embed + + +logger = get_logger(__name__) + + +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_deepspeed: 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, + ) + + # Optimizer creation + supported_optimizers = ["adam", "adamw", "prodigy"] + 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 and optimizer_name not in ["adam", "adamw"]: + logger.warning( + f"use_8bit_adam is ignored when optimizer is not set to 'Adam' or 'AdamW'. Optimizer was set to {optimizer_name}." + ) + + 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": + optimizer_class = bnb.optim.AdamW8bit if use_8bit else torch.optim.AdamW + + optimizer = optimizer_class( + params_to_optimize, + betas=(beta1, beta2), + eps=epsilon, + weight_decay=weight_decay, + ) + + elif optimizer_name == "adam": + optimizer_class = bnb.optim.Adam8bit if use_8bit else torch.optim.Adam + + optimizer = optimizer_class( + params_to_optimize, + 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" + ) + + optimizer = optimizer_class( + params_to_optimize, + lr=learning_rate, + betas=(beta1, beta2), + beta3=beta3, + weight_decay=weight_decay, + eps=epsilon, + decouple=prodigy_decouple, + use_bias_correction=prodigy_use_bias_correction, + safeguard_warmup=prodigy_safeguard_warmup, + ) + + return optimizer + + +def get_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 + + +# Similar to diffusers.pipelines.hunyuandit.pipeline_hunyuandit.get_resize_crop_region_for_grid +def get_resize_crop_region_for_grid(src, tgt_width, tgt_height): + tw = tgt_width + th = tgt_height + h, w = src + r = h / w + if r > (th / tw): + resize_height = th + resize_width = int(round(th / h * w)) + else: + resize_width = tw + resize_height = int(round(tw / w * h)) + + crop_top = int(round((th - resize_height) / 2.0)) + crop_left = int(round((tw - resize_width) / 2.0)) + + return (crop_top, crop_left), (crop_top + resize_height, crop_left + resize_width) + + +def prepare_rotary_positional_embeddings( + height: int, + width: int, + num_frames: int, + vae_scale_factor_spatial: int = 8, + patch_size: int = 2, + attention_head_dim: int = 64, + device: Optional[torch.device] = None, + base_height: int = 480, + base_width: int = 720, +) -> Tuple[torch.Tensor, torch.Tensor]: + grid_height = height // (vae_scale_factor_spatial * patch_size) + grid_width = width // (vae_scale_factor_spatial * patch_size) + base_size_width = base_width // (vae_scale_factor_spatial * patch_size) + base_size_height = base_height // (vae_scale_factor_spatial * patch_size) + + grid_crops_coords = get_resize_crop_region_for_grid((grid_height, grid_width), base_size_width, base_size_height) + freqs_cos, freqs_sin = get_3d_rotary_pos_embed( + embed_dim=attention_head_dim, + crops_coords=grid_crops_coords, + grid_size=(grid_height, grid_width), + temporal_size=num_frames, + ) + + freqs_cos = freqs_cos.to(device=device) + freqs_sin = freqs_sin.to(device=device) + return freqs_cos, freqs_sin + + +def reset_memory(device: Union[str, torch.device]) -> None: + gc.collect() + torch.cuda.empty_cache() + torch.cuda.reset_peak_memory_stats(device) + torch.cuda.reset_accumulated_memory_stats(device) + + +def print_memory(device: Union[str, torch.device]) -> None: + memory_allocated = torch.cuda.memory_allocated(device) / 1024**3 + max_memory_allocated = torch.cuda.max_memory_allocated(device) / 1024**3 + max_memory_reserved = torch.cuda.max_memory_reserved(device) / 1024**3 + print(f"{memory_allocated=:.3f} GB") + print(f"{max_memory_allocated=:.3f} GB") + print(f"{max_memory_reserved=:.3f} GB")