diff --git a/README.md b/README.md
index 093eb41..b150494 100644
--- a/README.md
+++ b/README.md
@@ -1,18 +1,67 @@
-# CogVideoX Factory
+# CogVideoX Factory 🧪
-## Introduction
+Fine-tune Cog family of video models for custom video generation under 24GB of GPU memory ⚡️📼
-This is a repos for CogVideoX Fine-tuning.
+
+
+ |
+
+
+## Quickstart
+Clone the repository and make sure the requirements are installed: `pip install -r requirements.txt`.
+
+Then download a dataset:
+
+```bash
+# install `huggingface_hub`
+huggingface-cli download \
+ --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset \
+ --local-dir video-dataset-disney
+```
+
+Then launch LoRA fine-tuning for text-to-video (modify the different hyperparameters, dataset root, and other configuration options as per your choice):
+
+```bash
+# For LoRA finetuning of the text-to-video CogVideoX models
+./train_text_to_video_lora.sh
+
+# For full finetuning of the text-to-video CogVideoX models
+./train_text_to_video_sft.sh
+
+# For LoRA finetuning of the image-to-video CogVideoX models
+./train_image_to_video_lora.sh
+```
+
+Assuming your LoRA is saved and pushed to the HF Hub, and named `my-awesome-name/my-awesome-lora`, we can now use the finetuned model for inference:
+
+```diff
+import torch
+from diffusers import CogVideoXPipeline
+from diffusers import export_to_video
+
+pipe = CogVideoXPipeline.from_pretrained(
+ "THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16
+).to("cuda")
++ pipe.load_lora_weights("my-awesome-name/my-awesome-lora", adapter_name=["cogvideox-lora"])
++ pipe.set_adapters(["cogvideox-lora"], [1.0])
+
+video = pipe("").frames[0]
+export_to_video(video, "output.mp4", fps=8)
+```
+
+**Note:** For Image-to-Video finetuning, you must install diffusers from [this](https://github.com/huggingface/diffusers/pull/9482) branch (which adds lora loading support in CogVideoX image-to-video) until it is merged.
+
+Below we provide additional sections detailing on more options explored in this repository. They all attempt to make fine-tuning for video models as accessible as possible by reducing memory requirements as much as possible.
## 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`.
+Assume you've specified `--data_root` as `/dataset`, and that this directory contains the files: `prompt.txt` and `videos.txt`.
-The `prompts.txt` file should contain line-separated prompts:
+The `prompt.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.
@@ -32,7 +81,7 @@ Overall, this is how your dataset would look like if you ran the `tree` command
```bash
/dataset
-├── prompts.txt
+├── prompt.txt
├── videos.txt
├── videos
├── videos/00000.mp4
@@ -40,7 +89,7 @@ Overall, this is how your dataset would look like if you ran the `tree` command
├── ...
```
-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.
+When using this format, the `--caption_column` must be `prompt.txt` and `--video_column` must be `videos.txt`. If you have your data stored in a CSV file instead, you can also specify `--dataset_file` as the path to CSV, and the `--caption_column` and `--video_column` as the actual column names in the CSV file. The [test_dataset](./tests/test_dataset.py) file contains some easy-to-understand examples for both formats.
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.
@@ -48,38 +97,160 @@ As an example, let's use [this](https://huggingface.co/datasets/Wild-Heart/Disne
huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset --local-dir video-dataset-disney
```
+This dataset is already prepared in the expected format and ready to use. However, using video datasets directly can lead to OOMs on smaller VRAM GPUs because it requires loading the [VAE](https://huggingface.co/THUDM/CogVideoX-5b/tree/main/vae) (to encode videos to latent space) and the massive [T5-XXL](https://huggingface.co/google/t5-v1_1-xxl/) text encoder. In order to lower these memory requirements, one can precompute the latents and embeddings using the `training/prepare_dataset.py` script.
+
+Fill in, or modify, the parameters in `prepare_dataset.sh` and execute it to obtain the precomputed latents and embeddings (make sure to specify `--save_tensors` to save precomputed artifacts). To use them during training, make sure to specify the `--load_tensors` flag, otherwise the videos will be used as-is and require loading the text encoder and VAE. The script also supports PyTorch DDP so that large datasets can be parallely encoded using multiple GPUs (modify the `NUM_GPUS` parameter).
+
## Training
-TODO
+We provide training script for both text-to-video and image-to-video generation which are compatible with the [CogVideoX family of models](https://huggingface.co/collections/THUDM/cogvideo-66c08e62f1685a3ade464cce). Training can be launched with one of the `train*.sh` scripts based on the task you'd like to train. Let's take text-to-video LoRA finetuning as an example.
-Take a look at `training/*.sh`
+- Configure environment variables according as per your choice:
-Note: Untested on MPS
+ ```bash
+ export TORCH_LOGS="+dynamo,recompiles,graph_breaks"
+ export TORCHDYNAMO_VERBOSE=1
+ export WANDB_MODE="offline"
+ export NCCL_P2P_DISABLE=1
+ export TORCH_NCCL_ENABLE_MONITORING=0
+ ```
+
+- Configure which GPUs to use for training: `GPU_IDS="0,1"`
+
+- Choose hyperparamters for training. Let's try to do a sweep on learning rate and optimizer type as an example:
+
+ ```bash
+ LEARNING_RATES=("1e-4" "1e-3")
+ LR_SCHEDULES=("cosine_with_restarts")
+ OPTIMIZERS=("adamw", "adam")
+ MAX_TRAIN_STEPS=("3000")
+ ```
+
+- Select which Accelerate configuration you would like to train with: `ACCELERATE_CONFIG_FILE="accelerate_configs/uncompiled_1.yaml"`. We provide some default configurations in the `accelerate_configs/` directory - single GPU uncompiled/compiled, 2x GPU DDP, DeepSpeed, etc. You can create your own config files with custom settings using `accelerate config --config_file my_config.yaml`.
+
+- Specify the absolute paths and columns/files for captions and videos.
+
+ ```bash
+ DATA_ROOT="/path/to/my/datasets/video-dataset-disney"
+ CAPTION_COLUMN="prompt.txt"
+ VIDEO_COLUMN="videos.txt"
+ ```
+
+- Launch experiments sweeping different hyperparameters:
+ ```
+ for learning_rate in "${LEARNING_RATES[@]}"; do
+ for lr_schedule in "${LR_SCHEDULES[@]}"; do
+ for optimizer in "${OPTIMIZERS[@]}"; do
+ for steps in "${MAX_TRAIN_STEPS[@]}"; do
+ output_dir="/path/to/my/models/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
+
+ cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox_text_to_video_lora.py \
+ --pretrained_model_name_or_path THUDM/CogVideoX-5b \
+ --data_root $DATA_ROOT \
+ --caption_column $CAPTION_COLUMN \
+ --video_column $VIDEO_COLUMN \
+ --id_token BW_STYLE \
+ --height_buckets 480 \
+ --width_buckets 720 \
+ --frame_buckets 49 \
+ --dataloader_num_workers 8 \
+ --pin_memory \
+ --validation_prompt \"BW_STYLE A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions:::BW_STYLE A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance\" \
+ --validation_prompt_separator ::: \
+ --num_validation_videos 1 \
+ --validation_epochs 10 \
+ --seed 42 \
+ --rank 128 \
+ --lora_alpha 128 \
+ --mixed_precision bf16 \
+ --output_dir $output_dir \
+ --max_num_frames 49 \
+ --train_batch_size 1 \
+ --max_train_steps $steps \
+ --checkpointing_steps 1000 \
+ --gradient_accumulation_steps 1 \
+ --gradient_checkpointing \
+ --learning_rate $learning_rate \
+ --lr_scheduler $lr_schedule \
+ --lr_warmup_steps 400 \
+ --lr_num_cycles 1 \
+ --enable_slicing \
+ --enable_tiling \
+ --optimizer $optimizer \
+ --beta1 0.9 \
+ --beta2 0.95 \
+ --weight_decay 0.001 \
+ --max_grad_norm 1.0 \
+ --allow_tf32 \
+ --report_to wandb \
+ --nccl_timeout 1800"
+
+ echo "Running command: $cmd"
+ eval $cmd
+ echo -ne "-------------------- Finished executing script --------------------\n\n"
+ done
+ done
+ done
+ done
+ ```
+
+ To understand what the different parameters mean, you could either take a look at the [args](./training/args.py) file or run the training script with `--help`.
+
+Note: Training scripts are untested on MPS, so performance and memory requirements can differ widely compared to the CUDA reports below.
## Memory requirements
Supported and verified memory optimizations for training include:
-- `CPUOffloadOptimizer` from [TorchAO](https://github.com/pytorch/ao). You can read about its capabilities and limitations [here](https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim#optimizer-cpu-offload). In short, it allows you to use the CPU for storing trainable parameters and gradients. This results in the optimizer step happening on the CPU, which requires a fast CPU optimizer, such as `torch.AdamW(fused=True)` or applying `torch.compile` on the optimizer step. Additionally, it is recommended to not `torch.compile` your model for training. Gradient clipping and accumulation is not supported yet either.
-- Low-bit optimizers from [bitsandbytes](https://huggingface.co/docs/bitsandbytes/optimizers). TODO: to test and make [TorchAO](https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim) ones work
-- TODO: DeepSpeed ZeRO
+
+- `CPUOffloadOptimizer` from [`torchao`](https://github.com/pytorch/ao). You can read about its capabilities and limitations [here](https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim#optimizer-cpu-offload). In short, it allows you to use the CPU for storing trainable parameters and gradients. This results in the optimizer step happening on the CPU, which requires a fast CPU optimizer, such as `torch.optim.AdamW(fused=True)` or applying `torch.compile` on the optimizer step. Additionally, it is recommended to not `torch.compile` your model for training. Gradient clipping and accumulation is not supported yet either.
+- Low-bit optimizers from [`bitsandbytes`](https://huggingface.co/docs/bitsandbytes/optimizers). TODO: to test and make [`torchao`](https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim) ones work
+- DeepSpeed Zero2: Since we rely on `accelerate`, follow [this guide](https://huggingface.co/docs/accelerate/en/usage_guides/deepspeed) to configure your `accelerate` installation to enable training with DeepSpeed Zero2 optimizations.
> [!IMPORTANT]
> The memory requirements are reported after running the `training/prepare_dataset.py`, which converts the videos and captions to latents and embeddings. During training, we directly load the latents and embeddings, and do not require the VAE or the T5 text encoder. However, if you perform validation/testing, these must be loaded and increase the amount of required memory. Not performing validation/testing saves a significant amount of memory, which can be used to focus solely on training if you're on smaller VRAM GPUs.
+>
+> If you choose to run validation/testing, you can save some memory on lower VRAM GPUs by specifying `--enable_model_cpu_offload`.
### LoRA finetuning
+> [!NOTE]
+> The memory requirements for image-to-video lora finetuning are similar to that of text-to-video on `THUDM/CogVideoX-5b`, so it hasn't been reported explicitly.
+>
+> Additionally, to prepare test images for I2V finetuning, you could either generate them on-the-fly by modifying the script, or extract some frames from your training data using:
+> `ffmpeg -i input.mp4 -frames:v 1 frame.png`,
+> or provide a URL to a valid and accessible image.
+
AdamW
+**Note:** Trying to run CogVideoX-5b without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
+
With `train_batch_size = 1`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
@@ -105,14 +276,13 @@ With `train_batch_size = 4`:
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 47.805 | 47.805 | 39.365 |
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 47.268 | 47.332 | 41.008 |
-> [!NOTE]
-> Trying to run CogVideoX-5b without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
-
AdamW (8-bit bitsandbytes)
+**Note:** Trying to run CogVideoX-5b without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
+
With `train_batch_size = 1`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
@@ -140,47 +310,11 @@ With `train_batch_size = 4`:
-
- AdamW (8-bit torchao)
-
-Currently, errors out with following stack-trace:
-
-```python
-Traceback (most recent call last):
- File "/raid/aryan/cogvideox-distillation/training/cogvideox_text_to_video_lora.py", line 915, in
- main(args)
- File "/raid/aryan/cogvideox-distillation/training/cogvideox_text_to_video_lora.py", line 719, in main
- optimizer.step()
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/accelerate/optimizer.py", line 159, in step
- self.scaler.step(self.optimizer, closure)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/amp/grad_scaler.py", line 457, in step
- retval = self._maybe_opt_step(optimizer, optimizer_state, *args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/amp/grad_scaler.py", line 352, in _maybe_opt_step
- retval = optimizer.step(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/accelerate/optimizer.py", line 214, in patched_step
- return method(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/optim/lr_scheduler.py", line 137, in wrapper
- return func.__get__(opt, opt.__class__)(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/optim/optimizer.py", line 487, in wrapper
- out = func(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 116, in decorate_context
- return func(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torchao/prototype/low_bit_optim/adam.py", line 87, in step
- raise RuntimeError(
-RuntimeError: lr was changed to a non-Tensor object. If you want to update lr, please use optim.param_groups[0]['lr'].fill_(new_lr)
-```
-
-
-
- AdamW (4-bit torchao)
-
-Same error as AdamW (8-bit torchao)
-
-
-
AdamW + CPUOffloadOptimizer (with gradient offloading)
+**Note:** Trying to run CogVideoX-5b without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
+
With `train_batch_size = 1`:
| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
@@ -206,57 +340,147 @@ With `train_batch_size = 4`:
| THUDM/CogVideoX-5b | 64 | True | 20.006 | 46.561 | 46.574 | 38.840 |
| THUDM/CogVideoX-5b | 256 | True | 20.771 | 47.268 | 47.332 | 39.623 |
-> [!NOTE]
-> Trying to run CogVideoX-5b without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
-
- AdamW (8-bit bitsandbytes) + CPUOffloadOptimizer (with gradient offloading)
+ DeepSpeed (AdamW + CPU/Parameter offloading)
-Currently, errors out with the following stack-trace:
+**Note:** Results are reported with `gradient_checkpointing` enabled, running on a 2x A100.
-```python
- File "/raid/aryan/cogvideox-distillation/training/cogvideox_text_to_video_lora.py", line 925, in
- main(args)
- File "/raid/aryan/cogvideox-distillation/training/cogvideox_text_to_video_lora.py", line 727, in main
- optimizer.step()
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 116, in decorate_context
- return func(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torchao/prototype/low_bit_optim/cpu_offload.py", line 87, in step
- self.optim_dict[p_cuda].step()
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/optim/optimizer.py", line 487, in wrapper
- out = func(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 116, in decorate_context
- return func(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/bitsandbytes/optim/optimizer.py", line 287, in step
- self.update_step(group, p, gindex, pindex)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/utils/_contextlib.py", line 116, in decorate_context
- return func(*args, **kwargs)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/bitsandbytes/optim/optimizer.py", line 546, in update_step
- F.optimizer_update_8bit_blockwise(
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/bitsandbytes/functional.py", line 1774, in optimizer_update_8bit_blockwise
- prev_device = pre_call(g.device)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/bitsandbytes/functional.py", line 463, in pre_call
- torch.cuda.set_device(device)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/cuda/__init__.py", line 476, in set_device
- device = _get_device_index(device)
- File "/raid/aryan/nightly-venv/lib/python3.10/site-packages/torch/cuda/_utils.py", line 34, in _get_device_index
- raise ValueError(f"Expected a cuda device, but got: {device}")
-ValueError: Expected a cuda device, but got: cpu
-```
+With `train_batch_size = 1`:
+
+| model | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | 13.141 | 13.141 | 21.070 | 24.602 |
+| THUDM/CogVideoX-5b | 20.170 | 20.170 | 28.662 | 38.957 |
+
+With `train_batch_size = 4`:
+
+| model | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | 13.141 | 19.854 | 20.836 | 24.709 |
+| THUDM/CogVideoX-5b | 20.170 | 40.635 | 40.699 | 39.027 |
### Full finetuning
> [!NOTE]
-> `memory_after_validation` is indicative of the peak memory required for training. This is because apart from the activations, parameters and gradients stored for training, you also need to load the vae and text encoder in memory and spend some memory to perform inference. In order to reduce total memory required to perform training, one can choose to not perform validation/testing as part of the training script.
+> The memory requirements for image-to-video full finetuning are similar to that of text-to-video on `THUDM/CogVideoX-5b`, so it hasn't been reported explicitly.
+>
+> Additionally, to prepare test images for I2V finetuning, you could either generate them on-the-fly by modifying the script, or extract some frames from your training data using:
+> `ffmpeg -i input.mp4 -frames:v 1 frame.png`,
+> or provide a URL to a valid and accessible image.
-- [ ] Make scripts compatible with DDP
+> [!NOTE]
+> Trying to run full finetuning without gradient checkpointing OOMs even on an A100 (80 GB), so the memory measurements have not been specified.
+
+
+ AdamW
+
+With `train_batch_size = 1`:
+
+| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | True | 16.396 | 33.934 | 43.848 | 37.520 |
+| THUDM/CogVideoX-5b | True | 30.061 | OOM | OOM | OOM |
+
+With `train_batch_size = 4`:
+
+| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | True | 16.396 | 38.281 | 48.341 | 37.544 |
+| THUDM/CogVideoX-5b | True | 30.061 | OOM | OOM | OOM |
+
+
+
+
+ AdamW (8-bit bitsandbytes)
+
+With `train_batch_size = 1`:
+
+| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | True | 16.396 | 16.447 | 27.555 | 27.156 |
+| THUDM/CogVideoX-5b | True | 30.061 | 52.826 | 58.570 | 49.541 |
+
+With `train_batch_size = 4`:
+
+| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | True | 16.396 | 27.930 | 27.990 | 27.326 |
+| THUDM/CogVideoX-5b | True | 16.396 | 66.648 | 66.705 | 48.828 |
+
+
+
+
+ AdamW + CPUOffloadOptimizer (with gradient offloading)
+
+With `train_batch_size = 1`:
+
+| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | True | 16.396 | 16.396 | 26.100 | 23.832 |
+| THUDM/CogVideoX-5b | True | 30.061 | 39.359 | 48.307 | 37.947 |
+
+With `train_batch_size = 4`:
+
+| model | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | True | 16.396 | 27.916 | 27.975 | 23.936 |
+| THUDM/CogVideoX-5b | True | 30.061 | 66.607 | 66.668 | 38.061 |
+
+
+
+
+ DeepSpeed (AdamW + CPU/Parameter offloading)
+
+**Note:** Results are reported with `gradient_checkpointing` enabled, running on a 2x A100.
+
+With `train_batch_size = 1`:
+
+| model | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | 13.111 | 13.111 | 20.328 | 23.867 |
+| THUDM/CogVideoX-5b | 19.762 | 19.998 | 27.697 | 38.018 |
+
+With `train_batch_size = 4`:
+
+| model | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | 13.111 | 21.188 | 21.254 | 23.869 |
+| THUDM/CogVideoX-5b | 19.762 | 43.465 | 43.531 | 38.082 |
+
+
+
+> [!NOTE]
+> - `memory_after_validation` is indicative of the peak memory required for training. This is because apart from the activations, parameters and gradients stored for training, you also need to load the vae and text encoder in memory and spend some memory to perform inference. In order to reduce total memory required to perform training, one can choose to not perform validation/testing as part of the training script.
+>
+> - `memory_before_validation` is the true indicator of the peak memory required for training if you choose to not perform validation/testing.
+
+
+
+## TODOs
+
+- [x] Make scripts compatible with DDP
- [ ] Make scripts compatible with FSDP
-- [ ] Make scripts compatible with DeepSpeed
+- [x] Make scripts compatible with DeepSpeed
+- [ ] vLLM-powered captioning script
+- [ ] Multi-resolution/frame support in `prepare_dataset.py`
+- [ ] Analyzing traces for potential speedups and removing as many syncs as possible
+- [ ] Support for QLoRA (priority), and other types of high usage LoRAs methods
- [x] Test scripts with memory-efficient optimizer from bitsandbytes
- [x] Test scripts with CPUOffloadOptimizer, etc.
-- [ ] Test scripts with torchao quantization, and low bit memory optimizers, etc.
-- [x] Make 5B lora finetuning work in under 24GB
+- [ ] Test scripts with torchao quantization, and low bit memory optimizers (Currently errors with AdamW (8/4-bit torchao))
+- [ ] Test scripts with AdamW (8-bit bitsandbytes) + CPUOffloadOptimizer (with gradient offloading) (Currently errors out)
+- [ ] [Sage Attention](https://github.com/thu-ml/SageAttention) (work with the authors to support backward pass, and optimize for A100)
+
+> [!IMPORTANT]
+> Since our goal is to make the scripts as memory-friendly as possible we don't guarantee multi-GPU training.
diff --git a/README_zh.md b/README_zh.md
index 250c7e3..3758d84 100644
--- a/README_zh.md
+++ b/README_zh.md
@@ -1,55 +1,24 @@
-# CogVideoX Factor 🧪
+# CogVideoX Factory
-在 24GB GPU 内存下微调 Cog 系列视频模型以生成自定义视频 ⚡️📼
+## 简介
-TODO:添加有趣的视频结果表
-
-## 快速开始
-
-确保已安装所需的依赖:`pip install -r requirements.txt`。
-
-然后下载数据集:
-
-```bash
-# 安装 `huggingface_hub`
-huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset --local-dir video-dataset-disney
-```
-
-然后启动文本到视频的 LoRA 微调:
-
-```bash
-TODO
-```
-
-我们现在可以使用训练好的模型进行推理:
-
-```python
-TODO
-```
-
-我们还可以使用 LoRA 微调 5B 版本:
-
-```python
-TODO
-```
-
-在下方的部分中,我们提供了有关更多选项的详细信息,这些选项旨在使视频模型的微调尽可能易于使用。
+这是用于 CogVideoX 微调的仓库。
## 数据集准备
-创建两个文件,一个文件包含逐行分隔的提示,另一个文件包含逐行分隔的视频数据路径(视频文件的路径必须相对于您在指定 `--data_root` 时传递的路径)。让我们通过一个示例来更好地理解这一点!
+创建两个文件,一个文件包含以换行符分隔的提示词,另一个文件包含以换行符分隔的视频数据路径(视频文件的路径必须相对于您在指定 `--data_root` 时传递的路径)。让我们通过一个例子来更好地理解这一点!
-假设您指定的 `--data_root` 为 `/dataset`,并且该目录包含以下文件:`prompts.txt` 和 `videos.txt`。
+假设您将 `--data_root` 指定为 `/dataset`,并且该目录包含文件:`prompts.txt` 和 `videos.txt`。
-`prompts.txt` 文件应包含逐行分隔的提示:
+`prompts.txt` 文件应包含以换行符分隔的提示词:
```
-一段黑白动画序列,主角是一只名为 Rabbity Ribfried 的兔子和一只拟人化的山羊,展示了它们在音乐与游戏环境中的互动演变。
-一段黑白动画序列,发生在船甲板上,主角是一只名为 Bully Bulldoger 的斗牛犬,展现了夸张的面部表情和肢体语言。角色从自信、专注逐渐转变为紧张与痛苦,展示了随着挑战出现的情感变化。船的内部在背景中保持静止,只有一些简单的细节,如钟声和敞开的门。角色的动态动作和不断变化的表情推动了叙事,没有摄像机运动来分散注意力。
+一段黑白动画序列,主角是一只名为 Rabbity Ribfried 的兔子和一只拟人化的山羊,在一个充满音乐和趣味的环境中,展示他们不断发展的互动。
+一段黑白动画序列,场景在船甲板上,主角是一只名为 Bully Bulldoger 的斗牛犬角色,展示了夸张的面部表情和肢体语言。角色从自信到专注,再到紧张和痛苦,展示了一系列情绪,随着它克服挑战。船的内部在背景中保持静止,只有简单的细节,如钟声和开着的门。角色的动态动作和变化的表情推动了故事的发展,没有镜头移动,确保观众专注于其不断变化的反应和肢体动作。
...
```
-`videos.txt` 文件应包含逐行分隔的视频文件路径。请注意,路径应相对于 `--data_root` 目录。
+`videos.txt` 文件应包含以换行符分隔的视频文件路径。请注意,路径应相对于 `--data_root` 目录。
```bash
videos/00000.mp4
@@ -57,7 +26,7 @@ videos/00001.mp4
...
```
-整体而言,如果在数据集根目录运行 `tree` 命令,您的数据集应如下所示:
+总体而言,如果您在数据集根目录运行 `tree` 命令,您的数据集应如下所示:
```bash
/dataset
@@ -69,54 +38,58 @@ videos/00001.mp4
├── ...
```
-使用此格式时,`--caption_column` 必须是 `prompts.txt`,`--video_column` 必须是 `videos.txt`。如果您将数据存储在 CSV 文件中,还可以指定 `--dataset_file` 为 CSV 的路径,`--caption_column` 和 `--video_column` 为 CSV 文件中的实际列名。
+使用此格式时,`--caption_column` 必须是 `prompts.txt`,`--video_column` 必须是 `videos.txt`。如果您的数据存储在 CSV 文件中,您也可以指定 `--dataset_file` 为 CSV 的路径,`--caption_column` 和 `--video_column` 为 CSV 文件中的实际列名。
-例如,让我们使用[这个](https://huggingface.co/datasets/Wild-Heart/Disney-VideoGeneration-Dataset) Disney 数据集进行微调。要下载,您可以使用 🤗 Hugging Face CLI。
+例如,让我们使用这个 [Disney 数据集](https://huggingface.co/datasets/Wild-Heart/Disney-VideoGeneration-Dataset) 进行微调。要下载,可以使用 🤗 Hugging Face CLI。
```bash
huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset --local-dir video-dataset-disney
```
-TODO:添加一个关于创建和使用预计算嵌入的部分。
-
## 训练
-我们提供了与 [Cog 系列模型](https://huggingface.co/collections/THUDM/cogvideo-66c08e62f1685a3ade464cce) 兼容的文本到视频和图像到视频生成的训练脚本。
+TODO
-查看 `*.sh` 文件
+请查看 `training/*.sh`
注意:未在 MPS 上测试
## 内存需求
-
+训练支持并验证的内存优化包括:
-支持和验证的内存优化训练选项包括:
-
-- [`torchao`](https://github.com/pytorch/ao) 中的 `CPUOffloadOptimizer`。您可以阅读它的能力和限制 [此处](https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim#optimizer-cpu-offload)。简而言之,它允许您使用 CPU 存储可训练的参数和梯度。这导致优化器步骤在 CPU 上进行,需要一个快速的 CPU 优化器,例如 `torch.optim.AdamW(fused=True)` 或在优化器步骤上应用 `torch.compile`。此外,建议不要将模型编译用于训练。梯度裁剪和积累尚不支持。
-- [`bitsandbytes`](https://huggingface.co/docs/bitsandbytes/optimizers) 中的低位优化器。TODO:测试并使 [`torchao`](https://github.com/pytorch/ao/tree/main/torchao/prototype/low_bit_optim) 工作
-- DeepSpeed Zero2:由于我们依赖 `accelerate`,请按照[本指南](https://huggingface.co/docs/accelerate/en/usage_guides/deepspeed) 配置 `accelerate` 以启用 DeepSpeed Zero2 优化。
-
-> [!IMPORTANT]
-> 内存需求是在运行 `training/prepare_dataset.py` 后报告的,它将视频和字幕转换为潜变量和嵌入。在训练过程中,我们直接加载潜变量和嵌入,而不需要 VAE 或 T5 文本编码器。但是,如果您执行验证/测试,则必须加载这些内容,并增加所需的内存量。不执行验证/测试可以节省大量内存,对于使用较小 VRAM 的 GPU,这可以用于专注于训练。
->
-> 如果您选择运行验证/测试,可以通过指定 `--enable_model_cpu_offload` 在较低 VRAM 的 GPU 上节省一些内存。
+- 来自 [TorchAO](https://github.com/pytorch/ao) 的 `CPUOffloadOptimizer`。
+- 来自 [bitsandbytes](https://huggingface.co/docs/bitsandbytes/optimizers) 的低位优化器。
### LoRA 微调
+
+ AdamW
+
+With `train_batch_size = 1`:
+
+| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | 16 | False | 12.945 | 43.764 | 46.918 | 24.234 |
+| THUDM/CogVideoX-2b | 16 | True | 12.945 | 12.945 | 21.121 | 24.234 |
+| THUDM/CogVideoX-2b | 64 | False | 13.035 | 44.314 | 47.469 | 24.469 |
+| THUDM/CogVideoX-2b | 64 | True | 13.036 | 13.035 | 21.564 | 24.500 |
+| THUDM/CogVideoX-2b | 256 | False | 13.095 | 45.826 | 48.990 | 25.543 |
+| THUDM/CogVideoX-2b | 256 | True | 13.094 | 13.095 | 22.344 | 25.537 |
+| THUDM/CogVideoX-5b | 16 | True | 19.742 | 19.742 | 28.746 | 38.123 |
+| THUDM/CogVideoX-5b | 64 | True | 20.006 | 20.818 | 30.338 | 38.738 |
+| THUDM/CogVideoX-5b | 256 | True | 20.771 | 22.119 | 31.939 | 41.537 |
+
+With `train_batch_size = 4`:
+
+| model | lora rank | gradient_checkpointing | memory_before_training | memory_before_validation | memory_after_validation | memory_after_testing |
+|:------------------:|:---------:|:----------------------:|:----------------------:|:------------------------:|:-----------------------:|:--------------------:|
+| THUDM/CogVideoX-2b | 16 | True | 12.945 | 21.803 | 21.814 | 24.322 |
+| THUDM/CogVideoX-2b | 64 | True | 13.035 | 22.254 | 22.254 | 24.572 |
+| THUDM/CogVideoX-2b | 256 | True | 13.094 | 22.020 | 22.033 | 25.574 |
+| THUDM/CogVideoX-5b | 16 | True | 19.742 | 46.492 | 46.492 | 38.197 |
+| THUDM/CogVideoX-5b | 64 | True | 20.006 | 47.805 | 47.805 | 39.365 |
+| THUDM/CogVideoX-5b | 256 | True | 20.771 | 47.268 | 47.332 | 41.008 |
+
> [!NOTE]
-> 图像到视频 LoRA 微调的内存需求与 `THUDM/CogVideoX-5b` 上的文本到视频类似,因此未明确报告。
->
-> 此外,要为 I2V 微调准备测试图像,您可以通过修改脚本动态生成它们,或使用以下命令从您的训练数据中提取一些帧:
-> `ffmpeg -i input.mp4 -frames:v 1 frame.png`,
-> 或提供一个有效且可访问的图像 URL。
-
-...
-
+>
\ No newline at end of file
diff --git a/accelerate_configs/deepspeed.yaml b/accelerate_configs/deepspeed.yaml
new file mode 100644
index 0000000..2827648
--- /dev/null
+++ b/accelerate_configs/deepspeed.yaml
@@ -0,0 +1,23 @@
+compute_environment: LOCAL_MACHINE
+debug: false
+deepspeed_config:
+ gradient_accumulation_steps: 1
+ gradient_clipping: 1.0
+ offload_optimizer_device: cpu
+ offload_param_device: cpu
+ zero3_init_flag: false
+ zero_stage: 2
+distributed_type: DEEPSPEED
+downcast_bf16: 'no'
+enable_cpu_affinity: false
+machine_rank: 0
+main_training_function: main
+mixed_precision: bf16
+num_machines: 1
+num_processes: 2
+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_2.yaml b/accelerate_configs/uncompiled_2.yaml
new file mode 100644
index 0000000..c5216da
--- /dev/null
+++ b/accelerate_configs/uncompiled_2.yaml
@@ -0,0 +1,17 @@
+compute_environment: LOCAL_MACHINE
+debug: false
+distributed_type: MULTI_GPU
+downcast_bf16: 'no'
+enable_cpu_affinity: false
+gpu_ids: 0,1
+machine_rank: 0
+main_training_function: main
+mixed_precision: bf16
+num_machines: 1
+num_processes: 2
+rdzv_backend: static
+same_network: true
+tpu_env: []
+tpu_use_cluster: false
+tpu_use_sudo: false
+use_cpu: false
diff --git a/assets/CogVideoX-LoRA.webm b/assets/CogVideoX-LoRA.webm
new file mode 100644
index 0000000..fd7a4e4
Binary files /dev/null and b/assets/CogVideoX-LoRA.webm differ
diff --git a/assets/lora_2b.png b/assets/lora_2b.png
new file mode 100644
index 0000000..44427b6
Binary files /dev/null and b/assets/lora_2b.png differ
diff --git a/assets/lora_5b.png b/assets/lora_5b.png
new file mode 100644
index 0000000..e0e1ea4
Binary files /dev/null and b/assets/lora_5b.png differ
diff --git a/assets/sft_2b.png b/assets/sft_2b.png
new file mode 100644
index 0000000..9340ef1
Binary files /dev/null and b/assets/sft_2b.png differ
diff --git a/assets/sft_5b.png b/assets/sft_5b.png
new file mode 100644
index 0000000..04509f3
Binary files /dev/null and b/assets/sft_5b.png differ
diff --git a/prepare_dataset.sh b/prepare_dataset.sh
index 84086a2..adacfa7 100755
--- a/prepare_dataset.sh
+++ b/prepare_dataset.sh
@@ -1,11 +1,14 @@
#!/bin/bash
-MODEL_ID="/share/official_pretrains/hf_home/CogVideoX-5b"
+MODEL_ID="THUDM/CogVideoX-2b"
-DATA_ROOT="/share/home/zyx/disney_cogvideox"
-CAPTION_COLUMN="prompts.txt"
+NUM_GPUS=8
+
+# For more details on the expected data format, please refer to the README.
+DATA_ROOT="/path/to/my/datasets/video-dataset" # This needs to be the path to the base directory where your videos are located.
+CAPTION_COLUMN="prompt.txt"
VIDEO_COLUMN="videos.txt"
-OUTPUT_DIR="/share/home/zyx/disney_cogvideox-encoded-multi"
+OUTPUT_DIR="/path/to/my/datasets/preprocessed-dataset"
HEIGHT=480
WIDTH=720
MAX_NUM_FRAMES=49
@@ -13,24 +16,31 @@ MAX_SEQUENCE_LENGTH=226
TARGET_FPS=8
BATCH_SIZE=1
DTYPE=fp32
-NUM_GPUS=8
-CMD="torchrun --nproc_per_node=$NUM_GPUS \
- 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 \
- --save_tensors"
+# To create a folder-style dataset structure without pre-encoding videos and captions
+# For Image-to-Video finetuning, make sure to pass `--save_image_latents`
+CMD_WITHOUT_PRE_ENCODING="\
+ torchrun --nproc_per_node=$NUM_GPUS \
+ 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"
\ No newline at end of file
+echo -ne "===== Finished running script =====\n"
diff --git a/train_image_to_video_lora.sh b/train_image_to_video_lora.sh
new file mode 100755
index 0000000..565f31a
--- /dev/null
+++ b/train_image_to_video_lora.sh
@@ -0,0 +1,82 @@
+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="0"
+
+# Training Configurations
+# Experiment with as many hyperparameters as you want!
+LEARNING_RATES=("1e-4" "1e-3")
+LR_SCHEDULES=("cosine_with_restarts")
+OPTIMIZERS=("adamw", "adam")
+MAX_TRAIN_STEPS=("3000")
+
+# Single GPU uncompiled training
+ACCELERATE_CONFIG_FILE="accelerate_configs/uncompiled_1.yaml"
+
+# Absolute path to where the data is located. Make sure to have read the README for how to prepare data.
+# This example assumes you downloaded an already prepared dataset from HF CLI as follows:
+# huggingface-cli download --repo-type dataset Wild-Heart/Disney-VideoGeneration-Dataset --local-dir /path/to/my/datasets/disney-dataset
+DATA_ROOT="/path/to/my/datasets/disney-dataset"
+CAPTION_COLUMN="prompt.txt"
+VIDEO_COLUMN="videos.txt"
+
+# Launch experiments with different hyperparameters
+for learning_rate in "${LEARNING_RATES[@]}"; do
+ for lr_schedule in "${LR_SCHEDULES[@]}"; do
+ for optimizer in "${OPTIMIZERS[@]}"; do
+ for steps in "${MAX_TRAIN_STEPS[@]}"; do
+ output_dir="/path/to/my/models/cogvideox-lora__optimizer_${optimizer}__steps_${steps}__lr-schedule_${lr_schedule}__learning-rate_${learning_rate}/"
+
+ cmd="accelerate launch --config_file $ACCELERATE_CONFIG_FILE --gpu_ids $GPU_IDS training/cogvideox_image_to_video_lora.py \
+ --pretrained_model_name_or_path THUDM/CogVideoX-5b-I2V \
+ --data_root $DATA_ROOT \
+ --caption_column $CAPTION_COLUMN \
+ --video_column $VIDEO_COLUMN \
+ --id_token BW_STYLE \
+ --height_buckets 480 \
+ --width_buckets 720 \
+ --frame_buckets 49 \
+ --dataloader_num_workers 8 \
+ --pin_memory \
+ --validation_prompt \"BW_STYLE A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions:::BW_STYLE A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance\" \
+ --validation_images \"/path/to/image1.png:::/path/to/image2.png\"
+ --validation_prompt_separator ::: \
+ --num_validation_videos 1 \
+ --validation_epochs 10 \
+ --seed 42 \
+ --rank 128 \
+ --lora_alpha 128 \
+ --mixed_precision bf16 \
+ --output_dir $output_dir \
+ --max_num_frames 49 \
+ --train_batch_size 1 \
+ --max_train_steps $steps \
+ --checkpointing_steps 1000 \
+ --gradient_accumulation_steps 1 \
+ --gradient_checkpointing \
+ --learning_rate $learning_rate \
+ --lr_scheduler $lr_schedule \
+ --lr_warmup_steps 400 \
+ --lr_num_cycles 1 \
+ --enable_slicing \
+ --enable_tiling \
+ --noised_image_dropout 0.05 \
+ --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_lora.sh b/train_text_to_video_lora.sh
index ab9013d..4aac214 100755
--- a/train_text_to_video_lora.sh
+++ b/train_text_to_video_lora.sh
@@ -43,6 +43,8 @@ for learning_rate in "${LEARNING_RATES[@]}"; do
--height_buckets 480 \
--width_buckets 720 \
--frame_buckets 49 \
+ --dataloader_num_workers 8 \
+ --pin_memory \
--validation_prompt \"BW_STYLE A black and white animated scene unfolds with an anthropomorphic goat surrounded by musical notes and symbols, suggesting a playful environment. Mickey Mouse appears, leaning forward in curiosity as the goat remains still. The goat then engages with Mickey, who bends down to converse or react. The dynamics shift as Mickey grabs the goat, potentially in surprise or playfulness, amidst a minimalistic background. The scene captures the evolving relationship between the two characters in a whimsical, animated setting, emphasizing their interactions and emotions:::BW_STYLE A panda, dressed in a small, red jacket and a tiny hat, sits on a wooden stool in a serene bamboo forest. The panda's fluffy paws strum a miniature acoustic guitar, producing soft, melodic tunes. Nearby, a few other pandas gather, watching curiously and some clapping in rhythm. Sunlight filters through the tall bamboo, casting a gentle glow on the scene. The panda's face is expressive, showing concentration and joy as it plays. The background includes a small, flowing stream and vibrant green foliage, enhancing the peaceful and magical atmosphere of this unique musical performance\" \
--validation_prompt_separator ::: \
--num_validation_videos 1 \
diff --git a/train_text_to_video_sft.sh b/train_text_to_video_sft.sh
index 8419a9b..9514154 100755
--- a/train_text_to_video_sft.sh
+++ b/train_text_to_video_sft.sh
@@ -38,6 +38,8 @@ for learning_rate in "${LEARNING_RATES[@]}"; do
--height_buckets 480 \
--width_buckets 720 \
--frame_buckets 49 \
+ --dataloader_num_workers 8 \
+ --pin_memory \
--validation_prompt \"Tom, the mischievous gray cat, is sprawled out on a vibrant red pillow, his body relaxed and his eyes half-closed, as if he's just woken up or is about to doze off. His white paws are stretched out in front of him, and his tail is casually draped over the edge of the pillow. The setting appears to be a cozy corner of a room, with a warm yellow wall in the background and a hint of a wooden floor. The scene captures a rare moment of tranquility for Tom, contrasting with his usual energetic and playful demeanor:::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 \
@@ -53,7 +55,7 @@ for learning_rate in "${LEARNING_RATES[@]}"; do
--gradient_checkpointing \
--learning_rate $learning_rate \
--lr_scheduler $lr_schedule \
- --lr_warmup_steps 200 \
+ --lr_warmup_steps 800 \
--lr_num_cycles 1 \
--enable_slicing \
--enable_tiling \
@@ -63,7 +65,7 @@ for learning_rate in "${LEARNING_RATES[@]}"; do
--weight_decay 0.001 \
--max_grad_norm 1.0 \
--allow_tf32 \
- --report_to wandb
+ --report_to wandb \
--nccl_timeout 1800"
echo "Running command: $cmd"
diff --git a/training/args.py b/training/args.py
index 5cf16e4..c20a471 100644
--- a/training/args.py
+++ b/training/args.py
@@ -96,6 +96,11 @@ def _get_dataset_args(parser: argparse.ArgumentParser) -> None:
default=0,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
+ parser.add_argument(
+ "--pin_memory",
+ action="store_true",
+ help="Whether or not to use the pinned memory setting in pytorch dataloader.",
+ )
def _get_validation_args(parser: argparse.ArgumentParser) -> None:
@@ -105,6 +110,12 @@ def _get_validation_args(parser: argparse.ArgumentParser) -> None:
default=None,
help="One or more prompt(s) that is used during validation to verify that the model is learning. Multiple validation prompts should be separated by the '--validation_prompt_seperator' string.",
)
+ parser.add_argument(
+ "--validation_images",
+ type=str,
+ default=None,
+ help="One or more image path(s)/URLs that is used during validation to verify that the model is learning. Multiple validation paths should be separated by the '--validation_prompt_seperator' string. These should correspond to the order of the validation prompts.",
+ )
parser.add_argument(
"--validation_prompt_separator",
type=str,
@@ -135,6 +146,12 @@ def _get_validation_args(parser: argparse.ArgumentParser) -> None:
default=False,
help="Whether or not to use the default cosine dynamic guidance schedule when sampling validation videos.",
)
+ parser.add_argument(
+ "--enable_model_cpu_offload",
+ action="store_true",
+ default=False,
+ help="Whether or not to enable model-wise CPU offloading when performing validation/testing to save memory.",
+ )
def _get_training_args(parser: argparse.ArgumentParser) -> None:
@@ -175,6 +192,12 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
default=720,
help="All input videos are resized to this width.",
)
+ parser.add_argument(
+ "--video_reshape_mode",
+ type=str,
+ default=None,
+ help="All input videos are reshaped to this mode. Choose between ['center', 'random', 'none']",
+ )
parser.add_argument("--fps", type=int, default=8, help="All input videos will be used at this FPS.")
parser.add_argument(
"--max_num_frames",
@@ -294,6 +317,12 @@ def _get_training_args(parser: argparse.ArgumentParser) -> None:
default=False,
help="Whether or not to use VAE tiling for saving memory.",
)
+ parser.add_argument(
+ "--noised_image_dropout",
+ type=float,
+ default=0.05,
+ help="Image condition dropout probability when finetuning image-to-video.",
+ )
def _get_optimizer_args(parser: argparse.ArgumentParser) -> None:
diff --git a/training/cogvideox_image_to_video_lora.py b/training/cogvideox_image_to_video_lora.py
new file mode 100644
index 0000000..885ae8a
--- /dev/null
+++ b/training/cogvideox_image_to_video_lora.py
@@ -0,0 +1,956 @@
+# 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 random
+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, DistributedType
+from accelerate.logging import get_logger
+from accelerate.utils import (
+ DistributedDataParallelKwargs,
+ InitProcessGroupKwargs,
+ ProjectConfiguration,
+ set_seed,
+)
+from diffusers import (
+ AutoencoderKLCogVideoX,
+ CogVideoXDPMScheduler,
+ CogVideoXImageToVideoPipeline,
+ 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, load_image
+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, VideoDatasetWithResizeAndRectangleCrop # 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 Finetune
+
+
+
+## Model description
+
+This is a lora finetune of the CogVideoX model `{base_model}`.
+
+The model was trained using [CogVideoX Factory](https://github.com/a-r-r-o-w/cogvideox-factory) - a repository containing memory-optimized training scripts for the CogVideoX family of models using [TorchAO](https://github.com/pytorch/ao) and [DeepSpeed](https://github.com/microsoft/DeepSpeed). The scripts were adopted from [CogVideoX Diffusers trainer](https://github.com/huggingface/diffusers/blob/main/examples/cogvideo/train_cogvideox_lora.py).
+
+## Download model
+
+[Download LoRA]({repo_id}/tree/main) in the Files & Versions tab.
+
+## Usage
+
+Requires the [🧨 Diffusers library](https://github.com/huggingface/diffusers) installed.
+
+```py
+import torch
+from diffusers import CogVideoXImageToVideoPipeline
+from diffusers.utils import export_to_video, load_image
+
+pipe = CogVideoXImageToVideoPipeline.from_pretrained("THUDM/CogVideoX-5b-I2V", 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])
+
+image = load_image("/path/to/image.png")
+video = pipe(image=image, prompt="{validation_prompt}", guidance_scale=6, use_dynamic_cfg=True).frames[0]
+export_to_video(video, "output.mp4", fps=8)
+```
+
+For more details, including weighting, merging and fusing LoRAs, check the [documentation](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters) on loading LoRAs in diffusers.
+
+## License
+
+Please adhere to the licensing terms as described [here](https://huggingface.co/THUDM/CogVideoX-5b-I2V/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",
+ "image-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: CogVideoXImageToVideoPipeline,
+ 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
+
+ # 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.bfloat16
+ else:
+ if accelerator.mixed_precision == "fp16":
+ weight_dtype = torch.float16
+ elif accelerator.mixed_precision == "bf16":
+ weight_dtype = torch.bfloat16
+
+ if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16:
+ # due to pytorch#99272, MPS does not yet support bfloat16.
+ raise ValueError(
+ "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead."
+ )
+
+ 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()
+
+ CogVideoXImageToVideoPipeline.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 = CogVideoXImageToVideoPipeline.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]
+ num_trainable_parameters = sum(param.numel() for model in params_to_optimize for param in model["params"])
+
+ use_deepspeed_optimizer = (
+ accelerator.state.deepspeed_plugin is not None
+ and "optimizer" in accelerator.state.deepspeed_plugin.deepspeed_config
+ )
+ use_deepspeed_scheduler = (
+ accelerator.state.deepspeed_plugin is not None
+ and "scheduler" in accelerator.state.deepspeed_plugin.deepspeed_config
+ )
+
+ optimizer = get_optimizer(
+ params_to_optimize=params_to_optimize,
+ optimizer_name=args.optimizer,
+ learning_rate=args.learning_rate,
+ beta1=args.beta1,
+ beta2=args.beta2,
+ beta3=args.beta3,
+ epsilon=args.epsilon,
+ weight_decay=args.weight_decay,
+ prodigy_decouple=args.prodigy_decouple,
+ prodigy_use_bias_correction=args.prodigy_use_bias_correction,
+ prodigy_safeguard_warmup=args.prodigy_safeguard_warmup,
+ use_8bit=args.use_8bit,
+ use_4bit=args.use_4bit,
+ use_torchao=args.use_torchao,
+ use_deepspeed=use_deepspeed_optimizer,
+ use_cpu_offload_optimizer=args.use_cpu_offload_optimizer,
+ offload_gradients=args.offload_gradients,
+ )
+
+ # Dataset and DataLoader
+ dataset_init_kwargs = {
+ "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,
+ "image_to_video": True,
+ }
+ if args.video_reshape_mode is None:
+ train_dataset = VideoDatasetWithResizing(**dataset_init_kwargs)
+ else:
+ train_dataset = VideoDatasetWithResizeAndRectangleCrop(
+ video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs
+ )
+
+ def collate_fn(data):
+ prompts = [x["prompt"] for x in data[0]]
+
+ if args.load_tensors:
+ prompts = torch.stack(prompts).to(dtype=weight_dtype, non_blocking=True)
+
+ images = [x["image"] for x in data[0]]
+ images = torch.stack(images).to(dtype=weight_dtype, non_blocking=True)
+
+ videos = [x["video"] for x in data[0]]
+ videos = torch.stack(videos).to(dtype=weight_dtype, non_blocking=True)
+
+ return {
+ "images": images,
+ "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,
+ num_workers=args.dataloader_num_workers,
+ pin_memory=args.pin_memory,
+ )
+
+ # 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 args.use_cpu_offload_optimizer:
+ lr_scheduler = None
+ accelerator.print(
+ "CPU Offload Optimizer cannot be used with DeepSpeed or builtin PyTorch LR Schedulers. If "
+ "you are training with those settings, they will be ignored."
+ )
+ else:
+ if use_deepspeed_scheduler:
+ from accelerate.utils import DummyScheduler
+
+ lr_scheduler = DummyScheduler(
+ name=args.lr_scheduler,
+ optimizer=optimizer,
+ total_num_steps=args.max_train_steps * accelerator.num_processes,
+ num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes,
+ )
+ else:
+ lr_scheduler = get_scheduler(
+ args.lr_scheduler,
+ optimizer=optimizer,
+ num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes,
+ num_training_steps=args.max_train_steps * accelerator.num_processes,
+ num_cycles=args.lr_num_cycles,
+ power=args.lr_power,
+ )
+
+ # Prepare everything with our `accelerator`.
+ transformer, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
+ transformer, optimizer, train_dataloader, lr_scheduler
+ )
+
+ # 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
+
+ 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)
+
+ alphas_cumprod = scheduler.alphas_cumprod.to(accelerator.device, dtype=torch.float32)
+
+ 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):
+ images = batch["images"].to(accelerator.device, non_blocking=True)
+ videos = batch["videos"].to(accelerator.device, non_blocking=True)
+ prompts = batch["prompts"]
+
+ # Encode videos
+ if not args.load_tensors:
+ image_noise_sigma = torch.normal(
+ mean=-3.0, std=0.5, size=(images.size(0),), device=accelerator.device, dtype=weight_dtype
+ )
+ image_noise_sigma = torch.exp(image_noise_sigma)
+ noisy_images = images + torch.randn_like(images) * image_noise_sigma[:, None, None, None, None]
+ image_latent_dist = vae.encode(noisy_images).latent_dist
+
+ videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
+ latent_dist = vae.encode(videos).latent_dist
+ else:
+ image_latent_dist = DiagonalGaussianDistribution(images)
+ latent_dist = DiagonalGaussianDistribution(videos)
+
+ image_latents = image_latent_dist.sample() * VAE_SCALING_FACTOR
+ image_latents = image_latents.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
+ image_latents = image_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
+
+ video_latents = latent_dist.sample() * VAE_SCALING_FACTOR
+ video_latents = video_latents.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
+ video_latents = video_latents.to(memory_format=torch.contiguous_format, dtype=weight_dtype)
+
+ padding_shape = (video_latents.shape[0], video_latents.shape[1] - 1, *video_latents.shape[2:])
+ latent_padding = image_latents.new_zeros(padding_shape)
+ image_latents = torch.cat([image_latents, latent_padding], dim=1)
+
+ if random.random() < args.noised_image_dropout:
+ image_latents = torch.zeros_like(image_latents)
+
+ # 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.to(dtype=weight_dtype)
+
+ # Sample noise that will be added to the latents
+ noise = torch.randn_like(video_latents)
+ batch_size, num_frames, num_channels, height, width = video_latents.shape
+
+ # Sample a random timestep for each image
+ timesteps = torch.randint(
+ 0,
+ scheduler.config.num_train_timesteps,
+ (batch_size,),
+ dtype=torch.int64,
+ device=accelerator.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_video_latents = scheduler.add_noise(video_latents, noise, timesteps)
+ noisy_model_input = torch.cat([noisy_video_latents, image_latents], dim=2)
+
+ # 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_video_latents, timesteps)
+
+ weights = 1 / (1 - alphas_cumprod[timesteps])
+ while len(weights.shape) < len(model_pred.shape):
+ weights = weights.unsqueeze(-1)
+
+ target = video_latents
+
+ 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()
+
+ if not args.use_cpu_offload_optimizer:
+ lr_scheduler.step()
+
+ # Checks if the accelerator has performed an optimization step behind the scenes
+ if accelerator.sync_gradients:
+ progress_bar.update(1)
+ global_step += 1
+
+ if accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED:
+ if global_step % args.checkpointing_steps == 0:
+ # _before_ saving state, check if this save would set us over the `checkpoints_total_limit`
+ if args.checkpoints_total_limit is not None:
+ checkpoints = os.listdir(args.output_dir)
+ checkpoints = [d for d in checkpoints if d.startswith("checkpoint")]
+ checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1]))
+
+ # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints
+ if len(checkpoints) >= args.checkpoints_total_limit:
+ num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1
+ removing_checkpoints = checkpoints[0:num_to_remove]
+
+ logger.info(
+ f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints"
+ )
+ logger.info(f"Removing checkpoints: {', '.join(removing_checkpoints)}")
+
+ for removing_checkpoint in removing_checkpoints:
+ removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint)
+ shutil.rmtree(removing_checkpoint)
+
+ save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
+ accelerator.save_state(save_path)
+ logger.info(f"Saved state to {save_path}")
+
+ last_lr = lr_scheduler.get_last_lr()[0] if lr_scheduler is not None else args.learning_rate
+ logs = {
+ "loss": loss.detach().item(),
+ "lr": last_lr,
+ "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 = CogVideoXImageToVideoPipeline.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()
+ if args.enable_model_cpu_offload:
+ pipe.enable_model_cpu_offload()
+
+ validation_prompts = args.validation_prompt.split(args.validation_prompt_separator)
+ validation_images = args.validation_images.split(args.validation_prompt_separator)
+ for validation_image, validation_prompt in zip(validation_images, validation_prompts):
+ pipeline_args = {
+ "image": load_image(validation_image),
+ "prompt": validation_prompt,
+ "guidance_scale": args.guidance_scale,
+ "use_dynamic_cfg": args.use_dynamic_cfg,
+ "height": args.height,
+ "width": args.width,
+ "max_sequence_length": model_config.max_text_seq_length,
+ }
+
+ 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)
+
+ CogVideoXImageToVideoPipeline.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 = CogVideoXImageToVideoPipeline.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()
+ if args.enable_model_cpu_offload:
+ pipe.enable_model_cpu_offload()
+
+ # Load LoRA weights
+ lora_scaling = args.lora_alpha / args.rank
+ pipe.load_lora_weights(args.output_dir, adapter_name="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)
+ validation_images = args.validation_images.split(args.validation_prompt_separator)
+ for validation_image, validation_prompt in zip(validation_images, validation_prompts):
+ pipeline_args = {
+ "image": load_image(validation_image),
+ "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_lora.py b/training/cogvideox_text_to_video_lora.py
index 2d258b1..514fbf4 100644
--- a/training/cogvideox_text_to_video_lora.py
+++ b/training/cogvideox_text_to_video_lora.py
@@ -25,7 +25,8 @@ from typing import Any, Dict
import diffusers
import torch
import transformers
-from accelerate import Accelerator
+import wandb
+from accelerate import Accelerator, DistributedType
from accelerate.logging import get_logger
from accelerate.utils import (
DistributedDataParallelKwargs,
@@ -51,11 +52,9 @@ from torch.utils.data import DataLoader
from tqdm.auto import tqdm
from transformers import AutoTokenizer, T5EncoderModel
-import wandb
-
from args import get_args # isort:skip
-from dataset import BucketSampler, VideoDatasetWithResizing # isort:skip
+from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # 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
@@ -83,30 +82,31 @@ def save_model_card(
)
model_description = f"""
-# CogVideoX LoRA - {repo_id}
+# CogVideoX LoRA Finetune
## Model description
-These are {repo_id} LoRA weights for {base_model}.
+This is a lora finetune of the CogVideoX model `{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.
+The model was trained using [CogVideoX Factory](https://github.com/a-r-r-o-w/cogvideox-factory) - a repository containing memory-optimized training scripts for the CogVideoX family of models using [TorchAO](https://github.com/pytorch/ao) and [DeepSpeed](https://github.com/microsoft/DeepSpeed). The scripts were adopted from [CogVideoX Diffusers trainer](https://github.com/huggingface/diffusers/blob/main/examples/cogvideo/train_cogvideox_lora.py).
## Download model
-[Download the *.safetensors LoRA]({repo_id}/tree/main) in the Files & versions tab.
+[Download LoRA]({repo_id}/tree/main) in the Files & Versions tab.
-## Use it with the [🧨 diffusers library](https://github.com/huggingface/diffusers)
+## Usage
+
+Requires the [🧨 Diffusers library](https://github.com/huggingface/diffusers) installed.
```py
-from diffusers import CogVideoXPipeline
import torch
+from diffusers import CogVideoXPipeline
+from diffusers.utils import export_to_video
pipe = CogVideoXPipeline.from_pretrained("THUDM/CogVideoX-5b", torch_dtype=torch.bfloat16).to("cuda")
-pipe.load_lora_weights("{repo_id}", weight_name="pytorch_lora_weights.safetensors", adapter_name=["cogvideox-lora"])
+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.
@@ -115,9 +115,10 @@ pipe.load_lora_weights("{repo_id}", weight_name="pytorch_lora_weights.safetensor
pipe.set_adapters(["cogvideox-lora"], [32 / 64])
video = pipe("{validation_prompt}", guidance_scale=6, use_dynamic_cfg=True).frames[0]
+export_to_video(video, "output.mp4", fps=8)
```
-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)
+For more details, including weighting, merging and fusing LoRAs, check the [documentation](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters) on loading LoRAs in diffusers.
## License
@@ -316,7 +317,7 @@ def main(args):
"bf16" in accelerator.state.deepspeed_plugin.deepspeed_config
and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"]
):
- weight_dtype = torch.float16
+ weight_dtype = torch.bfloat16
else:
if accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
@@ -437,7 +438,7 @@ def main(args):
)
use_deepspeed_scheduler = (
accelerator.state.deepspeed_plugin is not None
- and "scheduler" not in accelerator.state.deepspeed_plugin.deepspeed_config
+ and "scheduler" in accelerator.state.deepspeed_plugin.deepspeed_config
)
optimizer = get_optimizer(
@@ -461,46 +462,34 @@ def main(args):
)
# 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,
- )
+ dataset_init_kwargs = {
+ "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,
+ }
+ if args.video_reshape_mode is None:
+ train_dataset = VideoDatasetWithResizing(**dataset_init_kwargs)
+ else:
+ train_dataset = VideoDatasetWithResizeAndRectangleCrop(
+ video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs
+ )
- def collate_fn_without_pre_encoding(data):
+ def collate_fn(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)
+ if args.load_tensors:
+ prompts = torch.stack(prompts).to(dtype=weight_dtype, non_blocking=True)
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()
+ videos = torch.stack(videos).to(dtype=weight_dtype, non_blocking=True)
return {
"videos": videos,
@@ -511,8 +500,9 @@ def main(args):
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,
+ collate_fn=collate_fn,
num_workers=args.dataloader_num_workers,
+ pin_memory=args.pin_memory,
)
# Scheduler and math around the number of training steps.
@@ -637,9 +627,21 @@ def main(args):
models_to_accumulate = [transformer]
with accelerator.accumulate(models_to_accumulate):
- model_input = batch["videos"]
+ videos = batch["videos"].to(accelerator.device, non_blocking=True)
prompts = batch["prompts"]
+ # Encode videos
+ if not args.load_tensors:
+ videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
+ latent_dist = vae.encode(videos).latent_dist
+ else:
+ latent_dist = DiagonalGaussianDistribution(videos)
+
+ 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, dtype=weight_dtype)
+ model_input = videos
+
# Encode prompts
if not args.load_tensors:
prompt_embeds = compute_prompt_embeddings(
@@ -652,7 +654,7 @@ def main(args):
requires_grad=False,
)
else:
- prompt_embeds = prompts
+ prompt_embeds = prompts.to(dtype=weight_dtype)
# Sample noise that will be added to the latents
noise = torch.randn_like(model_input)
@@ -727,7 +729,7 @@ def main(args):
progress_bar.update(1)
global_step += 1
- if accelerator.is_main_process:
+ if accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED:
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:
@@ -785,6 +787,8 @@ def main(args):
pipe.vae.enable_slicing()
if args.enable_tiling:
pipe.vae.enable_tiling()
+ if args.enable_model_cpu_offload:
+ pipe.enable_model_cpu_offload()
validation_prompts = args.validation_prompt.split(args.validation_prompt_separator)
for validation_prompt in validation_prompts:
@@ -794,6 +798,7 @@ def main(args):
"use_dynamic_cfg": args.use_dynamic_cfg,
"height": args.height,
"width": args.width,
+ "max_sequence_length": model_config.max_text_seq_length,
}
log_validation(
@@ -859,6 +864,8 @@ def main(args):
pipe.vae.enable_slicing()
if args.enable_tiling:
pipe.vae.enable_tiling()
+ if args.enable_model_cpu_offload:
+ pipe.enable_model_cpu_offload()
# Load LoRA weights
lora_scaling = args.lora_alpha / args.rank
diff --git a/training/cogvideox_text_to_video_sft.py b/training/cogvideox_text_to_video_sft.py
index b10fd6e..52f3a45 100644
--- a/training/cogvideox_text_to_video_sft.py
+++ b/training/cogvideox_text_to_video_sft.py
@@ -25,7 +25,8 @@ from typing import Any, Dict
import diffusers
import torch
import transformers
-from accelerate import Accelerator
+import wandb
+from accelerate import Accelerator, DistributedType
from accelerate.logging import get_logger
from accelerate.utils import (
DistributedDataParallelKwargs,
@@ -50,11 +51,9 @@ from torch.utils.data import DataLoader
from tqdm.auto import tqdm
from transformers import AutoTokenizer, T5EncoderModel
-import wandb
-
from args import get_args # isort:skip
-from dataset import BucketSampler, VideoDatasetWithResizing # isort:skip
+from dataset import BucketSampler, VideoDatasetWithResizing, VideoDatasetWithResizeAndRectangleCrop # 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
@@ -81,7 +80,42 @@ def save_model_card(
}
)
- model_description = """TODO"""
+ model_description = f"""
+# CogVideoX Full Finetune
+
+
+
+## Model description
+
+This is a full finetune of the CogVideoX model `{base_model}`.
+
+The model was trained using [CogVideoX Factory](https://github.com/a-r-r-o-w/cogvideox-factory) - a repository containing memory-optimized training scripts for the CogVideoX family of models using [TorchAO](https://github.com/pytorch/ao) and [DeepSpeed](https://github.com/microsoft/DeepSpeed). The scripts were adopted from [CogVideoX Diffusers trainer](https://github.com/huggingface/diffusers/blob/main/examples/cogvideo/train_cogvideox_lora.py).
+
+## Download model
+
+[Download LoRA]({repo_id}/tree/main) in the Files & Versions tab.
+
+## Usage
+
+Requires the [🧨 Diffusers library](https://github.com/huggingface/diffusers) installed.
+
+```py
+import torch
+from diffusers import CogVideoXPipeline
+from diffusers.utils import export_to_video
+
+pipe = CogVideoXPipeline.from_pretrained("{repo_id}", torch_dtype=torch.bfloat16).to("cuda")
+
+video = pipe("{validation_prompt}", guidance_scale=6, use_dynamic_cfg=True).frames[0]
+export_to_video(video, "output.mp4", fps=8)
+```
+
+For more details, checkout the [documentation](https://huggingface.co/docs/diffusers/main/en/api/pipelines/cogvideox) for CogVideoX.
+
+## 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,
@@ -272,7 +306,7 @@ def main(args):
"bf16" in accelerator.state.deepspeed_plugin.deepspeed_config
and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"]
):
- weight_dtype = torch.float16
+ weight_dtype = torch.bfloat16
else:
if accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
@@ -368,7 +402,7 @@ def main(args):
)
use_deepspeed_scheduler = (
accelerator.state.deepspeed_plugin is not None
- and "scheduler" not in accelerator.state.deepspeed_plugin.deepspeed_config
+ and "scheduler" in accelerator.state.deepspeed_plugin.deepspeed_config
)
optimizer = get_optimizer(
@@ -392,46 +426,34 @@ def main(args):
)
# 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,
- )
+ dataset_init_kwargs = {
+ "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,
+ }
+ if args.video_reshape_mode is None:
+ train_dataset = VideoDatasetWithResizing(**dataset_init_kwargs)
+ else:
+ train_dataset = VideoDatasetWithResizeAndRectangleCrop(
+ video_reshape_mode=args.video_reshape_mode, **dataset_init_kwargs
+ )
- def collate_fn_without_pre_encoding(data):
+ def collate_fn(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)
+ if args.load_tensors:
+ prompts = torch.stack(prompts).to(dtype=weight_dtype, non_blocking=True)
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()
+ videos = torch.stack(videos).to(dtype=weight_dtype, non_blocking=True)
return {
"videos": videos,
@@ -442,8 +464,9 @@ def main(args):
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,
+ collate_fn=collate_fn,
num_workers=args.dataloader_num_workers,
+ pin_memory=args.pin_memory,
)
# Scheduler and math around the number of training steps.
@@ -568,9 +591,21 @@ def main(args):
models_to_accumulate = [transformer]
with accelerator.accumulate(models_to_accumulate):
- model_input = batch["videos"]
+ videos = batch["videos"].to(accelerator.device, non_blocking=True)
prompts = batch["prompts"]
+ # Encode videos
+ if not args.load_tensors:
+ videos = videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
+ latent_dist = vae.encode(videos).latent_dist
+ else:
+ latent_dist = DiagonalGaussianDistribution(videos)
+
+ 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, dtype=weight_dtype)
+ model_input = videos
+
# Encode prompts
if not args.load_tensors:
prompt_embeds = compute_prompt_embeddings(
@@ -583,7 +618,7 @@ def main(args):
requires_grad=False,
)
else:
- prompt_embeds = prompts
+ prompt_embeds = prompts.to(dtype=weight_dtype)
# Sample noise that will be added to the latents
noise = torch.randn_like(model_input)
@@ -658,7 +693,7 @@ def main(args):
progress_bar.update(1)
global_step += 1
- if accelerator.is_main_process:
+ if accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED:
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:
@@ -716,6 +751,8 @@ def main(args):
pipe.vae.enable_slicing()
if args.enable_tiling:
pipe.vae.enable_tiling()
+ if args.enable_model_cpu_offload:
+ pipe.enable_model_cpu_offload()
validation_prompts = args.validation_prompt.split(args.validation_prompt_separator)
for validation_prompt in validation_prompts:
@@ -725,6 +762,7 @@ def main(args):
"use_dynamic_cfg": args.use_dynamic_cfg,
"height": args.height,
"width": args.width,
+ "max_sequence_length": model_config.max_text_seq_length,
}
log_validation(
@@ -791,6 +829,8 @@ def main(args):
pipe.vae.enable_slicing()
if args.enable_tiling:
pipe.vae.enable_tiling()
+ if args.enable_model_cpu_offload:
+ pipe.enable_model_cpu_offload()
# Run inference
validation_outputs = []
diff --git a/training/dataset.py b/training/dataset.py
index fe125c9..2a6ac6c 100644
--- a/training/dataset.py
+++ b/training/dataset.py
@@ -2,11 +2,14 @@ import random
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
+import numpy as np
import pandas as pd
import torch
+import torchvision.transforms as TT
from accelerate.logging import get_logger
from torch.utils.data import Dataset, Sampler
from torchvision import transforms
+from torchvision.transforms import InterpolationMode
from torchvision.transforms.functional import resize
@@ -37,6 +40,7 @@ class VideoDataset(Dataset):
frame_buckets: List[int] = None,
load_tensors: bool = False,
random_flip: Optional[float] = None,
+ image_to_video: bool = False,
) -> None:
super().__init__()
@@ -51,6 +55,7 @@ class VideoDataset(Dataset):
self.frame_buckets = frame_buckets or FRAME_BUCKETS
self.load_tensors = load_tensors
self.random_flip = random_flip
+ self.image_to_video = image_to_video
self.resolutions = [
(f, h, w) for h in self.height_buckets for w in self.width_buckets for f in self.frame_buckets
@@ -104,23 +109,24 @@ class VideoDataset(Dataset):
return index
if self.load_tensors:
- latents, prompt_embeds = self._preprocess_video(self.video_paths[index])
+ image_latents, video_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)
+ latent_num_frames = video_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
+ height = video_latents.size(2) * 8
+ width = video_latents.size(3) * 8
return {
"prompt": prompt_embeds,
- "video": latents,
+ "image": image_latents,
+ "video": video_latents,
"video_metadata": {
"num_frames": num_frames,
"height": height,
@@ -128,10 +134,11 @@ class VideoDataset(Dataset):
},
}
else:
- video, _ = self._preprocess_video(self.video_paths[index])
+ image, video, _ = self._preprocess_video(self.video_paths[index])
return {
"prompt": self.id_token + self.prompts[index],
+ "image": image,
"video": video,
"video_metadata": {
"num_frames": video.shape[0],
@@ -204,7 +211,9 @@ class VideoDataset(Dataset):
frames = frames.permute(0, 3, 1, 2).contiguous()
frames = torch.stack([self.video_transforms(frame) for frame in frames], dim=0)
- return frames, None
+ image = frames[:1].clone() if self.image_to_video else None
+
+ return image, frames, None
def _load_preprocessed_latents_and_embeds(self, path: Path) -> Tuple[torch.Tensor, torch.Tensor]:
filename_without_ext = path.name.split(".")[0]
@@ -212,28 +221,34 @@ class VideoDataset(Dataset):
# The current path is something like: /a/b/c/d/videos/00001.mp4
# We need to reach: /a/b/c/d/latents/00001.pt
+ images_path = path.parent.parent.joinpath("image_latents")
latents_path = path.parent.parent.joinpath("latents")
embeds_path = path.parent.parent.joinpath("embeddings")
- if not latents_path.exists() or not embeds_path.exists():
+ if not latents_path.exists() or not embeds_path.exists() or (self.image_to_video and not images_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`."
+ 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`. Additionally, if you're training image-to-video, it is expected that an `image_latents` folder is also present."
)
+ if self.image_to_video:
+ image_filepath = images_path.joinpath(pt_filename)
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():
+ if self.image_to_video:
+ image_filepath = image_filepath.as_posix()
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`."
)
+ images = torch.load(image_filepath, map_location="cpu", weights_only=True) if self.image_to_video else None
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
+ return images, latents, embeds
class VideoDatasetWithResizing(VideoDataset):
@@ -258,9 +273,76 @@ class VideoDatasetWithResizing(VideoDataset):
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
+
+ image = frames[:1].clone() if self.image_to_video else None
+
+ return image, 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 VideoDatasetWithResizeAndRectangleCrop(VideoDataset):
+ def __init__(self, video_reshape_mode: str = "center", *args, **kwargs) -> None:
+ super().__init__(*args, **kwargs)
+ self.video_reshape_mode = video_reshape_mode
+
+ def _resize_for_rectangle_crop(self, arr, image_size):
+ reshape_mode = self.video_reshape_mode
+ if arr.shape[3] / arr.shape[2] > image_size[1] / image_size[0]:
+ arr = resize(
+ arr,
+ size=[image_size[0], int(arr.shape[3] * image_size[0] / arr.shape[2])],
+ interpolation=InterpolationMode.BICUBIC,
+ )
+ else:
+ arr = resize(
+ arr,
+ size=[int(arr.shape[2] * image_size[1] / arr.shape[3]), image_size[1]],
+ interpolation=InterpolationMode.BICUBIC,
+ )
+
+ h, w = arr.shape[2], arr.shape[3]
+ arr = arr.squeeze(0)
+
+ delta_h = h - image_size[0]
+ delta_w = w - image_size[1]
+
+ if reshape_mode == "random" or reshape_mode == "none":
+ top = np.random.randint(0, delta_h + 1)
+ left = np.random.randint(0, delta_w + 1)
+ elif reshape_mode == "center":
+ top, left = delta_h // 2, delta_w // 2
+ else:
+ raise NotImplementedError
+ arr = TT.functional.crop(arr, top=top, left=left, height=image_size[0], width=image_size[1])
+ return arr
+
+ def _preprocess_video(self, path: Path) -> torch.Tensor:
+ 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 = self._resize_for_rectangle_crop(frames, nearest_res)
+ frames = torch.stack([self.video_transforms(frame) for frame in frames_resized], dim=0)
+
+ image = frames[:1].clone() if self.image_to_video else None
+
+ return image, 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))
diff --git a/training/prepare_dataset.py b/training/prepare_dataset.py
index 4e54fab..48e5b81 100644
--- a/training/prepare_dataset.py
+++ b/training/prepare_dataset.py
@@ -7,14 +7,19 @@ import pathlib
import traceback
from typing import Any, Dict, List, Optional, Tuple, Union
+import numpy as np
import pandas as pd
import torch
import torch.distributed as dist
+import torchvision.transforms as TT
from diffusers import AutoencoderKLCogVideoX
from diffusers.utils import export_to_video, get_logger
from torchvision import transforms
-from transformers import T5EncoderModel, T5Tokenizer
+from torchvision.transforms import InterpolationMode
+from torchvision.transforms.functional import resize
from tqdm import tqdm
+from transformers import T5EncoderModel, T5Tokenizer
+
import decord # isort:skip
@@ -53,6 +58,11 @@ def get_args() -> Dict[str, Any]:
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(
+ "--save_image_latents",
+ action="store_true",
+ help="Whether or not to encode and store image latents, which are required for image-to-video finetuning. The image latents are the first frame of input videos encoded with the VAE.",
+ )
parser.add_argument(
"--output_dir",
type=str,
@@ -147,13 +157,51 @@ def load_dataset_from_csv(
return prompts, video_paths
+def resize_for_rectangle_crop(arr, height, width, reshape_mode):
+ image_size = height, width
+ if arr.shape[3] / arr.shape[2] > image_size[1] / image_size[0]:
+ arr = resize(
+ arr,
+ size=[image_size[0], int(arr.shape[3] * image_size[0] / arr.shape[2])],
+ interpolation=InterpolationMode.BICUBIC,
+ )
+ else:
+ arr = resize(
+ arr,
+ size=[int(arr.shape[2] * image_size[1] / arr.shape[3]), image_size[1]],
+ interpolation=InterpolationMode.BICUBIC,
+ )
+
+ h, w = arr.shape[2], arr.shape[3]
+ arr = arr.squeeze(0)
+
+ delta_h = h - image_size[0]
+ delta_w = w - image_size[1]
+
+ if reshape_mode == "random" or reshape_mode == "none":
+ top = np.random.randint(0, delta_h + 1)
+ left = np.random.randint(0, delta_w + 1)
+ elif reshape_mode == "center":
+ top, left = delta_h // 2, delta_w // 2
+ else:
+ raise NotImplementedError
+ arr = TT.functional.crop(arr, top=top, left=left, height=image_size[0], width=image_size[1])
+ return arr
+
+
def load_and_preprocess_video(
- path: pathlib.Path, height: int, width: int, max_num_frames: int, video_transforms, num_threads: int = 0
+ path: pathlib.Path,
+ height: int,
+ width: int,
+ max_num_frames: int,
+ video_transforms,
+ num_threads: int = 0,
+ video_reshape_mode: str = "center",
) -> Optional[torch.Tensor]:
frames = None
try:
- video_reader = decord.VideoReader(uri=path.as_posix(), height=height, width=width, num_threads=num_threads)
+ video_reader = decord.VideoReader(uri=path.as_posix(), num_threads=num_threads)
video_num_frames = len(video_reader)
if video_num_frames < max_num_frames:
@@ -166,6 +214,7 @@ def load_and_preprocess_video(
frames: torch.Tensor = video_reader.get_batch(indices)
frames = frames[:max_num_frames].float()
frames = frames.permute(0, 3, 1, 2).contiguous()
+ frames = resize_for_rectangle_crop(frames, height, width, video_reshape_mode)
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()}`")
@@ -270,7 +319,11 @@ def compute_prompt_embeddings(
def save_videos(
- videos: torch.Tensor, video_paths: List[pathlib.Path], prompts: List[str], output_dir: pathlib.Path, target_fps: int = 8
+ videos: torch.Tensor,
+ video_paths: List[pathlib.Path],
+ prompts: List[str],
+ output_dir: pathlib.Path,
+ target_fps: int = 8,
) -> None:
assert videos.size(0) == len(video_paths)
@@ -291,50 +344,61 @@ def save_videos(
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:
+ with open(output_dir.joinpath("videos.txt").as_posix(), "a", encoding="utf-8") as file:
for video_path in video_paths:
file.write(f"videos/{video_path.name}\n")
- with open(output_dir.joinpath("prompts.txt").as_posix(), "w", encoding="utf-8") as file:
+ with open(output_dir.joinpath("prompt.txt").as_posix(), "a", encoding="utf-8") as file:
for prompt in prompts:
file.write(f"{prompt}\n")
def save_latents_and_embeddings(
+ image_latents: torch.Tensor,
latents: torch.Tensor,
prompt_embeds: torch.Tensor,
video_paths: List[pathlib.Path],
prompts: List[str],
output_dir: pathlib.Path,
+ save_image_latents: bool = False,
) -> None:
assert latents.size(0) == prompt_embeds.size(0)
assert latents.size(0) == len(video_paths)
assert prompt_embeds.size(0) == len(prompts)
+ if save_image_latents:
+ assert image_latents.size(0) == latents.size(0)
+ else:
+ image_latents = [None] * latents.size(0)
+ image_latents_dir = output_dir.joinpath("image_latents")
latents_dir = output_dir.joinpath("latents")
embeds_dir = output_dir.joinpath("embeddings")
output_dir.mkdir(parents=True, exist_ok=True)
+ image_latents_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):
+ for image_latent, latent, embed, video_path in zip(image_latents, latents, prompt_embeds, video_paths):
+ image_latent = image_latent.clone()
latent = latent.clone()
embed = embed.clone()
filename_without_ext = video_path.stem
+ image_latent_filename = image_latents_dir.joinpath(f"{filename_without_ext}.pt")
latent_filename = latents_dir.joinpath(f"{filename_without_ext}.pt")
embed_filename = embeds_dir.joinpath(f"{filename_without_ext}.pt")
+ torch.save(image_latent, image_latent_filename)
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:
+ with open(output_dir.joinpath("videos.txt").as_posix(), "a", encoding="utf-8") as file:
for video_path in video_paths:
file.write(f"videos/{video_path.name}\n")
- with open(output_dir.joinpath("prompts.txt").as_posix(), "w", encoding="utf-8") as file:
+ with open(output_dir.joinpath("prompt.txt").as_posix(), "a", encoding="utf-8") as file:
for prompt in prompts:
file.write(f"{prompt}\n")
@@ -344,8 +408,8 @@ def main():
args = get_args()
# Initialize distributed processing
- if 'LOCAL_RANK' in os.environ:
- local_rank = int(os.environ['LOCAL_RANK'])
+ if "LOCAL_RANK" in os.environ:
+ local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl")
world_size = dist.get_world_size()
@@ -447,7 +511,9 @@ def main():
)
prompt_embeds_list.append(prompt_embeds.to("cpu"))
- prompt_embeds = torch.cat(prompt_embeds_list)
+ prompt_embeds = None
+ if len(prompt_embeds_list) > 0:
+ prompt_embeds = torch.cat(prompt_embeds_list)
del tokenizer, text_encoder
gc.collect()
@@ -462,7 +528,8 @@ def main():
if args.use_tiling:
vae.enable_tiling()
- encoded_videos = []
+ encoded_videos_list = []
+ encoded_images_list = []
if rank == 0:
iterator = tqdm(range(0, len(video_paths_usable), args.batch_size), desc="Encoding videos")
@@ -476,15 +543,33 @@ def main():
batch_videos = batch_videos.to(device)
batch_videos = batch_videos.permute(0, 2, 1, 3, 4) # [B, C, F, H, W]
+ if args.save_image_latents:
+ batch_images = batch_videos[:, :, :1].clone()
+
if args.use_slicing:
encoded_slices = [vae._encode(video_slice) for video_slice in batch_videos.split(1)]
encoded_video = torch.cat(encoded_slices)
+ encoded_videos_list.append(encoded_video.to("cpu"))
+
+ if args.save_image_latents:
+ encoded_slices = [vae._encode(image_slice) for image_slice in batch_images.split(1)]
+ encoded_image = torch.cat(encoded_slices)
+ encoded_images_list.append(encoded_image.to("cpu"))
else:
encoded_video = vae._encode(batch_videos)
+ encoded_videos_list.append(encoded_video.to("cpu"))
- encoded_videos.append(encoded_video.to("cpu"))
+ if args.save_image_latents:
+ encoded_image = vae._encode(batch_images)
+ encoded_images_list.append(encoded_image.to("cpu"))
- encoded_videos = torch.cat(encoded_videos)
+ encoded_videos = None
+ if len(encoded_videos_list) > 0:
+ encoded_videos = torch.cat(encoded_videos_list)
+
+ encoded_images = None
+ if len(encoded_images_list) > 0:
+ encoded_images = torch.cat(encoded_images_list)
del vae
gc.collect()
@@ -495,9 +580,17 @@ def main():
if world_size > 1:
dist.barrier()
- save_latents_and_embeddings(
- encoded_videos, prompt_embeds, video_paths_usable, prompts_usable, pathlib.Path(args.output_dir)
- )
+ if prompt_embeds is not None:
+ assert encoded_videos is not None
+ save_latents_and_embeddings(
+ encoded_images,
+ encoded_videos,
+ prompt_embeds,
+ video_paths_usable,
+ prompts_usable,
+ pathlib.Path(args.output_dir),
+ args.save_image_latents,
+ )
# Finalize distributed processing
if world_size > 1:
@@ -514,4 +607,4 @@ if __name__ == "__main__":
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()
\ No newline at end of file
+ main()