diff --git a/README.md b/README.md index 5a47af3..81916ad 100644 --- a/README.md +++ b/README.md @@ -108,6 +108,8 @@ Supported and verified memory optimizations for training include: > [!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_offloading`. ### LoRA finetuning @@ -307,7 +309,64 @@ With `train_batch_size = 4`: ### 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. +> 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) @@ -331,7 +390,10 @@ With `train_batch_size = 4`:
-- [ ] Make scripts compatible with DDP +> [!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. + +- [x] Make scripts compatible with DDP - [ ] Make scripts compatible with FSDP - [x] Make scripts compatible with DeepSpeed - [x] Test scripts with memory-efficient optimizer from bitsandbytes diff --git a/training/args.py b/training/args.py index f987eff..03f4f4c 100644 --- a/training/args.py +++ b/training/args.py @@ -140,6 +140,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_offloading", + 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: diff --git a/training/cogvideox_text_to_video_lora.py b/training/cogvideox_text_to_video_lora.py index 3347c37..3dbf684 100644 --- a/training/cogvideox_text_to_video_lora.py +++ b/training/cogvideox_text_to_video_lora.py @@ -779,6 +779,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: @@ -853,6 +855,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 4fe58af..b98b270 100644 --- a/training/cogvideox_text_to_video_sft.py +++ b/training/cogvideox_text_to_video_sft.py @@ -710,6 +710,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: @@ -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() # Run inference validation_outputs = []