diff --git a/README.md b/README.md index dc615ec..d1d4f39 100644 --- a/README.md +++ b/README.md @@ -12,7 +12,8 @@ FineTrainers is a work-in-progress library to support (accessible) training of v ## News -- 🔥 **2024-12-20**: Support for T2V LoRA finetuning of [CogVideoX](https://huggingface.co/docs/diffusers/main/api/pipelines/cogvideox) added! +- 🔥 **2024-01-13**: Support for T2V full-finetuning added! Thanks to @ArEnSc for taking up the initiative! +- 🔥 **2024-01-03**: Support for T2V LoRA finetuning of [CogVideoX](https://huggingface.co/docs/diffusers/main/api/pipelines/cogvideox) added! - 🔥 **2024-12-20**: Support for T2V LoRA finetuning of [Hunyuan Video](https://huggingface.co/docs/diffusers/main/api/pipelines/hunyuan_video) added! We would like to thank @SHYuanBest for his work on a training script [here](https://github.com/huggingface/diffusers/pull/10254). - 🔥 **2024-12-18**: Support for T2V LoRA finetuning of [LTX Video](https://huggingface.co/docs/diffusers/main/api/pipelines/ltx_video) added! @@ -137,17 +138,16 @@ For inference, refer [here](./docs/training/ltx_video.md#inference). For docs re
-| **Model Name** | **Tasks** | **Min. GPU VRAM** | -|:---:|:---:|:---:| -| [LTX-Video](./docs/training/ltx_video.md) | Text-to-Video | 11 GB | -| [HunyuanVideo](./docs/training/hunyuan_video.md) | Text-to-Video | 42 GB | -| [CogVideoX](./docs/training/cogvideox.md) | Text-to-Video | 12GB* | +| **Model Name** | **Tasks** | **Min. LoRA VRAM*** | **Min. Full Finetuning VRAM^** | +|:------------------------------------------------:|:-------------:|:----------------------------------:|:---------------------------------------------:| +| [LTX-Video](./docs/training/ltx_video.md) | Text-to-Video | 11 GB | 21 GB | +| [HunyuanVideo](./docs/training/hunyuan_video.md) | Text-to-Video | 42 GB | OOM | +| [CogVideoX-5b](./docs/training/cogvideox.md) | Text-to-Video | 21 GB | 53 GB |
-*Noted for the 5B variant. - -Note that the memory consumption in the table is reported with most of the options, discussed in [docs/training/optimizations](./docs/training/optimization.md), enabled. +*Noted for training-only, no validation, at resolution `49x512x768`, rank 128, with pre-computation, using fp8 weights & gradient checkpointing. Pre-computation of conditions and latents may require higher limits (but typically under 16 GB).
+^Noted for training-only, no validation, at resolution `49x512x768`, with pre-computation, using bf16 weights & gradient checkpointing. If you would like to use a custom dataset, refer to the dataset preparation guide [here](./docs/dataset/README.md). diff --git a/docs/training/README.md b/docs/training/README.md index c53c80d..6a109c3 100644 --- a/docs/training/README.md +++ b/docs/training/README.md @@ -1,8 +1,9 @@ -This directory contains the training-related specifications for all the models we support in `finetrainers`. Each model page has: +# FineTrainers training documentation -* an example training command -* inference example -* numbers on memory consumption +This directory contains the training-related specifications for all the models we support in `finetrainers`. Each model page has: +- an example training command +- inference example +- numbers on memory consumption By default, we don't include any validation-related arguments in the example training commands. To enable validation inference, one can pass: @@ -12,8 +13,13 @@ By default, we don't include any validation-related arguments in the example tra + --validation_steps 100 ``` -## Model-specific docs +Supported models: +- [CogVideoX](./cogvideox.md) +- [LTX-Video](./ltx_video.md) +- [HunyuanVideo](./hunyuan_video.md) -* [CogVideoX](./cogvideox.md) -* [LTX-Video](./ltx_video.md) -* [HunyuanVideo](./hunyuan_video.md) \ No newline at end of file +Supported training types: +- LoRA (`--training_type lora`) +- Full finetuning (`--training_type full-finetune`) + +Arguments for training are well-documented in the code. For more information, please run `python train.py --help`. diff --git a/docs/training/cogvideox.md b/docs/training/cogvideox.md index 3784786..3900d25 100644 --- a/docs/training/cogvideox.md +++ b/docs/training/cogvideox.md @@ -2,6 +2,8 @@ ## Training +For LoRA training, specify `--training_type lora`. For full finetuning, specify `--training_type full-finetune`. + ```bash #!/bin/bash export WANDB_MODE="offline" @@ -84,6 +86,8 @@ echo -ne "-------------------- Finished executing script --------------------\n\ ## Memory Usage +### LoRA + LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x480x720` resolutions, **with precomputation**: ``` @@ -109,6 +113,31 @@ Training configuration: { | after validation end | 11.145 | 28.324 | | after training end | 11.144 | 11.592 | +### Full finetuning + +``` +Training configuration: { + "trainable parameters": 5570283072, + "total samples": 1, + "train epochs": 2, + "train steps": 2, + "batches per device": 1, + "total batches observed per epoch": 1, + "train batch size": 1, + "gradient accumulation steps": 1 +} +``` + +| stage | memory_allocated | max_memory_reserved | +|:-----------------------------:|:-----------------:|:-------------------:| +| after precomputing conditions | 8.880 | 8.941 | +| after precomputing latents | 9.300 | 12.441 | +| before training start | 10.376 | 10.387 | +| after epoch 1 | 31.160 | 52.939 | +| before validation start | 31.161 | 52.939 | +| after validation end | 31.161 | 52.939 | +| after training end | 31.160 | 34.295 | + ## Supported checkpoints CogVideoX has multiple checkpoints as one can note [here](https://huggingface.co/collections/THUDM/cogvideo-66c08e62f1685a3ade464cce). The following checkpoints were tested with `finetrainers` and are known to be working: diff --git a/docs/training/hunyuan_video.md b/docs/training/hunyuan_video.md index e62657a..10ef2dd 100644 --- a/docs/training/hunyuan_video.md +++ b/docs/training/hunyuan_video.md @@ -2,6 +2,8 @@ ## Training +For LoRA training, specify `--training_type lora`. For full finetuning, specify `--training_type full-finetune`. + ```bash #!/bin/bash @@ -87,6 +89,8 @@ echo -ne "-------------------- Finished executing script --------------------\n\ ## Memory Usage +### LoRA + LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolutions, **without precomputation**: ``` @@ -139,6 +143,10 @@ Training configuration: { Note: requires about `47` GB of VRAM with validation. If validation is not performed, the memory usage is reduced to about `42` GB. +### Full finetuning + +Current, full finetuning is not supported for HunyuanVideo. It goes out of memory (OOM) for `49x512x768` resolutions. + ## Inference Assuming your LoRA is saved and pushed to the HF Hub, and named `my-awesome-name/my-awesome-lora`, we can now use the finetuned model for inference: diff --git a/docs/training/ltx_video.md b/docs/training/ltx_video.md index a25390c..f55f459 100644 --- a/docs/training/ltx_video.md +++ b/docs/training/ltx_video.md @@ -2,7 +2,7 @@ ## Training -Provided you have a dataset: +For LoRA training, specify `--training_type lora`. For full finetuning, specify `--training_type full-finetune`. ```bash #!/bin/bash @@ -88,6 +88,8 @@ echo -ne "-------------------- Finished executing script --------------------\n\ ## Memory Usage +### LoRA + LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolution, **without precomputation**: ``` @@ -140,6 +142,31 @@ Training configuration: { Note: requires about `17.5` GB of VRAM with precomputation. If validation is not performed, the memory usage is reduced to `11` GB. +### Full Finetuning + +``` +Training configuration: { + "trainable parameters": 1923385472, + "total samples": 1, + "train epochs": 10, + "train steps": 10, + "batches per device": 1, + "total batches observed per epoch": 1, + "train batch size": 1, + "gradient accumulation steps": 1 +} +``` + +| stage | memory_allocated | max_memory_reserved | +|:-----------------------------:|:----------------:|:-------------------:| +| after precomputing conditions | 8.89 | 8.937 | +| after precomputing latents | 9.701 | 11.615 | +| before training start | 3.583 | 4.025 | +| after epoch 1 | 10.769 | 20.357 | +| before validation start | 10.769 | 20.357 | +| after validation end | 10.769 | 28.332 | +| after training end | 10.769 | 12.904 | + ## Inference Assuming your LoRA is saved and pushed to the HF Hub, and named `my-awesome-name/my-awesome-lora`, we can now use the finetuned model for inference: diff --git a/finetrainers/args.py b/finetrainers/args.py index 92f9472..1f2cfd4 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -207,6 +207,8 @@ class Args: Perform validation every `n` training steps. enable_model_cpu_offload (`bool`, defaults to `False`): Whether or not to offload different modeling components to CPU during validation. + validation_frame_rate (`int`, defaults to `25`): + Frame rate to use for the validation videos. This value is defaulted to 25, as used in LTX Video pipeline. MISCELLANEOUS ARGUMENTS ----------------------- @@ -319,6 +321,7 @@ class Args: validation_every_n_epochs: Optional[int] = None validation_every_n_steps: Optional[int] = None enable_model_cpu_offload: bool = False + validation_frame_rate: int = 25 # Miscellaneous arguments tracker_name: str = "finetrainers" @@ -417,6 +420,7 @@ class Args: "validation_every_n_epochs": self.validation_every_n_epochs, "validation_every_n_steps": self.validation_every_n_steps, "enable_model_cpu_offload": self.enable_model_cpu_offload, + "validation_frame_rate": self.validation_frame_rate, }, "miscellaneous_arguments": { "tracker_name": self.tracker_name, @@ -460,6 +464,7 @@ def parse_arguments() -> Args: def validate_args(args: Args): + _validate_training_args(args) _validate_validation_args(args) @@ -678,8 +683,9 @@ def _add_training_arguments(parser: argparse.ArgumentParser) -> None: parser.add_argument( "--training_type", type=str, + choices=["lora", "full-finetune"], required=True, - help="Type of training to perform. Choose between ['lora']", + help="Type of training to perform. Choose between ['lora', 'full-finetune']", ) parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") parser.add_argument( @@ -713,7 +719,11 @@ def _add_training_arguments(parser: argparse.ArgumentParser) -> None: help="The lora_alpha to compute scaling factor (lora_alpha / rank) for LoRA matrices.", ) parser.add_argument( - "--target_modules", type=str, default="to_k to_q to_v to_out.0", nargs="+", help="The target modules for LoRA." + "--target_modules", + type=str, + default=["to_k", "to_q", "to_v", "to_out.0"], + nargs="+", + help="The target modules for LoRA.", ) parser.add_argument( "--gradient_accumulation_steps", @@ -890,6 +900,12 @@ def _add_validation_arguments(parser: argparse.ArgumentParser) -> None: default=None, help="Run validation every X training steps. Validation consists of running the validation prompt `args.num_validation_videos` times.", ) + parser.add_argument( + "--validation_frame_rate", + type=int, + default=25, + help="Frame rate to use for the validation videos.", + ) parser.add_argument( "--enable_model_cpu_offload", action="store_true", @@ -1085,6 +1101,7 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: result_args.validation_every_n_epochs = args.validation_epochs result_args.validation_every_n_steps = args.validation_steps result_args.enable_model_cpu_offload = args.enable_model_cpu_offload + result_args.validation_frame_rate = args.validation_frame_rate # Miscellaneous arguments result_args.tracker_name = args.tracker_name @@ -1100,6 +1117,15 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args: return result_args +def _validate_training_args(args: Args): + if args.training_type == "lora": + assert args.rank is not None, "Rank is required for LoRA training" + assert args.lora_alpha is not None, "LoRA alpha is required for LoRA training" + assert ( + args.target_modules is not None and len(args.target_modules) > 0 + ), "Target modules are required for LoRA training" + + def _validate_validation_args(args: Args): assert args.validation_prompts is not None, "Validation prompts are required for validation" if args.validation_images is not None: diff --git a/finetrainers/cogvideox/__init__.py b/finetrainers/cogvideox/__init__.py index 6a3f826..390479b 100644 --- a/finetrainers/cogvideox/__init__.py +++ b/finetrainers/cogvideox/__init__.py @@ -1 +1,2 @@ from .cogvideox_lora import COGVIDEOX_T2V_LORA_CONFIG +from .full_finetune import COGVIDEOX_T2V_FULL_FINETUNE_CONFIG diff --git a/finetrainers/cogvideox/cogvideox_lora.py b/finetrainers/cogvideox/cogvideox_lora.py index c3b754a..7dca3d0 100644 --- a/finetrainers/cogvideox/cogvideox_lora.py +++ b/finetrainers/cogvideox/cogvideox_lora.py @@ -311,6 +311,7 @@ def _pad_frames(latents: torch.Tensor, patch_size_t: int): return latents +# TODO(aryan): refactor into model specs for better re-use COGVIDEOX_T2V_LORA_CONFIG = { "pipeline_cls": CogVideoXPipeline, "load_condition_models": load_condition_models, diff --git a/finetrainers/cogvideox/full_finetune.py b/finetrainers/cogvideox/full_finetune.py new file mode 100644 index 0000000..f755981 --- /dev/null +++ b/finetrainers/cogvideox/full_finetune.py @@ -0,0 +1,32 @@ +from diffusers import CogVideoXPipeline + +from .cogvideox_lora import ( + calculate_noisy_latents, + collate_fn_t2v, + forward_pass, + initialize_pipeline, + load_condition_models, + load_diffusion_models, + load_latent_models, + post_latent_preparation, + prepare_conditions, + prepare_latents, + validation, +) + + +# TODO(aryan): refactor into model specs for better re-use +COGVIDEOX_T2V_FULL_FINETUNE_CONFIG = { + "pipeline_cls": CogVideoXPipeline, + "load_condition_models": load_condition_models, + "load_latent_models": load_latent_models, + "load_diffusion_models": load_diffusion_models, + "initialize_pipeline": initialize_pipeline, + "prepare_conditions": prepare_conditions, + "prepare_latents": prepare_latents, + "post_latent_preparation": post_latent_preparation, + "collate_fn": collate_fn_t2v, + "calculate_noisy_latents": calculate_noisy_latents, + "forward_pass": forward_pass, + "validation": validation, +} diff --git a/finetrainers/hunyuan_video/__init__.py b/finetrainers/hunyuan_video/__init__.py index f4e780d..e1fdafa 100644 --- a/finetrainers/hunyuan_video/__init__.py +++ b/finetrainers/hunyuan_video/__init__.py @@ -1 +1,2 @@ +from .full_finetune import HUNYUAN_VIDEO_T2V_FULL_FINETUNE_CONFIG from .hunyuan_video_lora import HUNYUAN_VIDEO_T2V_LORA_CONFIG diff --git a/finetrainers/hunyuan_video/full_finetune.py b/finetrainers/hunyuan_video/full_finetune.py new file mode 100644 index 0000000..36dd5cb --- /dev/null +++ b/finetrainers/hunyuan_video/full_finetune.py @@ -0,0 +1,30 @@ +from diffusers import HunyuanVideoPipeline + +from .hunyuan_video_lora import ( + collate_fn_t2v, + forward_pass, + initialize_pipeline, + load_condition_models, + load_diffusion_models, + load_latent_models, + post_latent_preparation, + prepare_conditions, + prepare_latents, + validation, +) + + +# TODO(aryan): refactor into model specs for better re-use +HUNYUAN_VIDEO_T2V_FULL_FINETUNE_CONFIG = { + "pipeline_cls": HunyuanVideoPipeline, + "load_condition_models": load_condition_models, + "load_latent_models": load_latent_models, + "load_diffusion_models": load_diffusion_models, + "initialize_pipeline": initialize_pipeline, + "prepare_conditions": prepare_conditions, + "prepare_latents": prepare_latents, + "post_latent_preparation": post_latent_preparation, + "collate_fn": collate_fn_t2v, + "forward_pass": forward_pass, + "validation": validation, +} diff --git a/finetrainers/hunyuan_video/hunyuan_video_lora.py b/finetrainers/hunyuan_video/hunyuan_video_lora.py index 9bfea53..ed9013c 100644 --- a/finetrainers/hunyuan_video/hunyuan_video_lora.py +++ b/finetrainers/hunyuan_video/hunyuan_video_lora.py @@ -345,6 +345,7 @@ def _get_clip_prompt_embeds( return {"pooled_prompt_embeds": prompt_embeds} +# TODO(aryan): refactor into model specs for better re-use HUNYUAN_VIDEO_T2V_LORA_CONFIG = { "pipeline_cls": HunyuanVideoPipeline, "load_condition_models": load_condition_models, diff --git a/finetrainers/ltx_video/__init__.py b/finetrainers/ltx_video/__init__.py index b583686..6d5d0f9 100644 --- a/finetrainers/ltx_video/__init__.py +++ b/finetrainers/ltx_video/__init__.py @@ -1 +1,2 @@ +from .full_finetune import LTX_VIDEO_T2V_FULL_FINETUNE_CONFIG from .ltx_video_lora import LTX_VIDEO_T2V_LORA_CONFIG diff --git a/finetrainers/ltx_video/full_finetune.py b/finetrainers/ltx_video/full_finetune.py new file mode 100644 index 0000000..9aa30ef --- /dev/null +++ b/finetrainers/ltx_video/full_finetune.py @@ -0,0 +1,30 @@ +from diffusers import LTXPipeline + +from .ltx_video_lora import ( + collate_fn_t2v, + forward_pass, + initialize_pipeline, + load_condition_models, + load_diffusion_models, + load_latent_models, + post_latent_preparation, + prepare_conditions, + prepare_latents, + validation, +) + + +# TODO(aryan): refactor into model specs for better re-use +LTX_VIDEO_T2V_FULL_FINETUNE_CONFIG = { + "pipeline_cls": LTXPipeline, + "load_condition_models": load_condition_models, + "load_latent_models": load_latent_models, + "load_diffusion_models": load_diffusion_models, + "initialize_pipeline": initialize_pipeline, + "prepare_conditions": prepare_conditions, + "prepare_latents": prepare_latents, + "post_latent_preparation": post_latent_preparation, + "collate_fn": collate_fn_t2v, + "forward_pass": forward_pass, + "validation": validation, +} diff --git a/finetrainers/ltx_video/ltx_video_lora.py b/finetrainers/ltx_video/ltx_video_lora.py index 0e1af9b..c5c1df2 100644 --- a/finetrainers/ltx_video/ltx_video_lora.py +++ b/finetrainers/ltx_video/ltx_video_lora.py @@ -225,7 +225,7 @@ def validation( height: Optional[int] = None, width: Optional[int] = None, num_frames: Optional[int] = None, - frame_rate: int = 25, + frame_rate: int = 24, num_videos_per_prompt: int = 1, generator: Optional[torch.Generator] = None, **kwargs, diff --git a/finetrainers/models.py b/finetrainers/models.py index c7d95ae..c24ab95 100644 --- a/finetrainers/models.py +++ b/finetrainers/models.py @@ -1,19 +1,22 @@ from typing import Any, Dict -from .cogvideox import COGVIDEOX_T2V_LORA_CONFIG -from .hunyuan_video import HUNYUAN_VIDEO_T2V_LORA_CONFIG -from .ltx_video import LTX_VIDEO_T2V_LORA_CONFIG +from .cogvideox import COGVIDEOX_T2V_FULL_FINETUNE_CONFIG, COGVIDEOX_T2V_LORA_CONFIG +from .hunyuan_video import HUNYUAN_VIDEO_T2V_FULL_FINETUNE_CONFIG, HUNYUAN_VIDEO_T2V_LORA_CONFIG +from .ltx_video import LTX_VIDEO_T2V_FULL_FINETUNE_CONFIG, LTX_VIDEO_T2V_LORA_CONFIG SUPPORTED_MODEL_CONFIGS = { "hunyuan_video": { "lora": HUNYUAN_VIDEO_T2V_LORA_CONFIG, + "full-finetune": HUNYUAN_VIDEO_T2V_FULL_FINETUNE_CONFIG, }, "ltx_video": { "lora": LTX_VIDEO_T2V_LORA_CONFIG, + "full-finetune": LTX_VIDEO_T2V_FULL_FINETUNE_CONFIG, }, "cogvideox": { "lora": COGVIDEOX_T2V_LORA_CONFIG, + "full-finetune": COGVIDEOX_T2V_FULL_FINETUNE_CONFIG, }, } diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index 8f2ebed..a63dd5c 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -5,7 +5,7 @@ import os import random from datetime import datetime, timedelta from pathlib import Path -from typing import Any, Dict +from typing import Any, Dict, List import diffusers import torch @@ -21,6 +21,7 @@ from accelerate.utils import ( gather_object, set_seed, ) +from diffusers import DiffusionPipeline from diffusers.configuration_utils import FrozenDict from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution from diffusers.optimization import get_scheduler @@ -242,16 +243,7 @@ class Trainer: condition_components = self.model_config["load_condition_models"](**self._get_load_components_kwargs()) self._set_components(condition_components) self._move_components_to_device() - - # TODO(aryan): refactor later. for now only lora is supported - components_to_disable_grads = [ - self.text_encoder, - self.text_encoder_2, - self.text_encoder_3, - ] - for component in components_to_disable_grads: - if component is not None: - component.requires_grad_(False) + self._disable_grad_for_components([self.text_encoder, self.text_encoder_2, self.text_encoder_3]) if self.args.caption_dropout_p > 0 and self.args.caption_dropout_technique == "empty": logger.warning( @@ -305,12 +297,7 @@ class Trainer: latent_components = self.model_config["load_latent_models"](**self._get_load_components_kwargs()) self._set_components(latent_components) self._move_components_to_device() - - # TODO(aryan): refactor later - components_to_disable_grads = [self.vae] - for component in components_to_disable_grads: - if component is not None: - component.requires_grad_(False) + self._disable_grad_for_components([self.vae]) if self.vae is not None: if self.args.enable_slicing: @@ -371,24 +358,22 @@ class Trainer: diffusion_components = self.model_config["load_diffusion_models"](**self._get_load_components_kwargs()) self._set_components(diffusion_components) - # TODO(aryan): refactor later. for now only lora is supported - components_to_disable_grads = [ - self.text_encoder, - self.text_encoder_2, - self.text_encoder_3, - self.transformer, - self.vae, - ] - for component in components_to_disable_grads: - if component is not None: - component.requires_grad_(False) + components = [self.text_encoder, self.text_encoder_2, self.text_encoder_3, self.vae] + self._disable_grad_for_components(components) + + if self.args.training_type == "full-finetune": + logger.info("Finetuning transformer with no additional parameters") + self._enable_grad_for_components([self.transformer]) + else: + logger.info("Finetuning transformer with PEFT parameters") + self._disable_grad_for_components([self.transformer]) # 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 = self._get_training_dtype(accelerator=self.state.accelerator) if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: - # due to pytorch#99272, MPS does not yet support 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." ) @@ -406,13 +391,16 @@ class Trainer: if self.args.gradient_checkpointing: self.transformer.enable_gradient_checkpointing() - transformer_lora_config = LoraConfig( - r=self.args.rank, - lora_alpha=self.args.lora_alpha, - init_lora_weights=True, - target_modules=self.args.target_modules, - ) - self.transformer.add_adapter(transformer_lora_config) + if self.args.training_type == "lora": + transformer_lora_config = LoraConfig( + r=self.args.rank, + lora_alpha=self.args.lora_alpha, + init_lora_weights=True, + target_modules=self.args.target_modules, + ) + self.transformer.add_adapter(transformer_lora_config) + else: + transformer_lora_config = None # Enable TF32 for faster training on Ampere GPUs: https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices if self.args.allow_tf32 and torch.cuda.is_available(): @@ -432,7 +420,8 @@ class Trainer: type(unwrap_model(self.state.accelerator, self.transformer)), ): model = unwrap_model(self.state.accelerator, model) - transformer_lora_layers_to_save = get_peft_model_state_dict(model) + if self.args.training_type == "lora": + transformer_lora_layers_to_save = get_peft_model_state_dict(model) else: raise ValueError(f"Unexpected save model: {model.__class__}") @@ -440,10 +429,18 @@ class Trainer: if weights: weights.pop() - self.model_config["pipeline_cls"].save_lora_weights( - output_dir, - transformer_lora_layers=transformer_lora_layers_to_save, - ) + if self.args.training_type == "lora": + self.model_config["pipeline_cls"].save_lora_weights( + output_dir, + transformer_lora_layers=transformer_lora_layers_to_save, + ) + else: + model.save_pretrained(os.path.join(output_dir, "transformer")) + + # In some cases, the scheduler needs to be loaded with specific config (e.g. in CogVideoX). Since we need + # to able to load all diffusion components from a specific checkpoint folder during validation, we need to + # ensure the scheduler config is serialized as well. + self.scheduler.save_pretrained(os.path.join(output_dir, "scheduler")) def load_model_hook(models, input_dir): if not self.state.accelerator.distributed_type == DistributedType.DEEPSPEED: @@ -459,33 +456,39 @@ class Trainer: f"Unexpected save model: {unwrap_model(self.state.accelerator, model).__class__}" ) else: - transformer_ = unwrap_model(self.state.accelerator, self.transformer).__class__.from_pretrained( - self.args.pretrained_model_name_or_path, subfolder="transformer" - ) - transformer_.add_adapter(transformer_lora_config) + transformer_cls_ = unwrap_model(self.state.accelerator, self.transformer).__class__ - lora_state_dict = self.model_config["pipeline_cls"].lora_state_dict(input_dir) - transformer_state_dict = { - f'{k.replace("transformer.", "")}': v - for k, v in lora_state_dict.items() - if k.startswith("transformer.") - } - incompatible_keys = set_peft_model_state_dict(transformer_, transformer_state_dict, adapter_name="default") - if incompatible_keys is not None: - # check only for unexpected keys - unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None) - if unexpected_keys: - logger.warning( - f"Loading adapter weights from state_dict led to unexpected keys not found in the model: " - f" {unexpected_keys}. " + if self.args.training_type == "lora": + transformer_ = transformer_cls_.from_pretrained( + self.args.pretrained_model_name_or_path, subfolder="transformer" ) + transformer_.add_adapter(transformer_lora_config) + lora_state_dict = self.model_config["pipeline_cls"].lora_state_dict(input_dir) + transformer_state_dict = { + f'{k.replace("transformer.", "")}': v + for k, v in lora_state_dict.items() + if k.startswith("transformer.") + } + incompatible_keys = set_peft_model_state_dict( + transformer_, transformer_state_dict, adapter_name="default" + ) + if incompatible_keys is not None: + # check only for unexpected keys + unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None) + if unexpected_keys: + logger.warning( + f"Loading adapter weights from state_dict led to unexpected keys not found in the model: " + f" {unexpected_keys}. " + ) - # Make sure the trainable params are in float32. This is again needed since the base models - # are in `weight_dtype`. More details: - # https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804 - if self.args.mixed_precision == "fp16": - # only upcast trainable parameters (LoRA) into fp32 - cast_training_params([transformer_]) + # Make sure the trainable params are in float32. This is again needed since the base models + # are in `weight_dtype`. More details: + # https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804 + if self.args.mixed_precision == "fp16" and self.args.training_type == "lora": + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params([transformer_], dtype=torch.float32) + else: + transformer_ = transformer_cls_.from_pretrained(os.path.join(input_dir, "transformer")) self.state.accelerator.register_save_state_pre_hook(save_model_hook) self.state.accelerator.register_load_state_pre_hook(load_model_hook) @@ -497,7 +500,7 @@ class Trainer: self.state.train_steps = self.args.train_steps # Make sure the trainable params are in float32 - if self.args.mixed_precision == "fp16": + if self.args.mixed_precision == "fp16" and self.args.training_type == "lora": # only upcast trainable parameters (LoRA) into fp32 cast_training_params([self.transformer], dtype=torch.float32) @@ -510,13 +513,13 @@ class Trainer: * self.state.accelerator.num_processes ) - transformer_lora_parameters = list(filter(lambda p: p.requires_grad, self.transformer.parameters())) + transformer_trainable_parameters = list(filter(lambda p: p.requires_grad, self.transformer.parameters())) transformer_parameters_with_lr = { - "params": transformer_lora_parameters, + "params": transformer_trainable_parameters, "lr": self.state.learning_rate, } params_to_optimize = [transformer_parameters_with_lr] - self.state.num_trainable_parameters = sum(p.numel() for p in transformer_lora_parameters) + self.state.num_trainable_parameters = sum(p.numel() for p in transformer_trainable_parameters) use_deepspeed_opt = ( self.state.accelerator.state.deepspeed_plugin is not None @@ -608,6 +611,12 @@ class Trainer: ) self.vae_config = FrozenDict(**vae_config) + # In some cases, the scheduler needs to be loaded with specific config (e.g. in CogVideoX). Since we need + # to able to load all diffusion components from a specific checkpoint folder during validation, we need to + # ensure the scheduler config is serialized as well. + if self.args.training_type == "full-finetune": + self.scheduler.save_pretrained(os.path.join(self.args.output_dir, "scheduler")) + self.state.train_batch_size = ( self.args.batch_size * self.state.accelerator.num_processes * self.args.gradient_accumulation_steps ) @@ -872,14 +881,17 @@ class Trainer: accelerator.wait_for_everyone() if accelerator.is_main_process: - # TODO: consider factoring this out when supporting other types of training algos. - self.transformer = unwrap_model(accelerator, self.transformer) - transformer_lora_layers = get_peft_model_state_dict(self.transformer) + transformer = unwrap_model(accelerator, self.transformer) - self.model_config["pipeline_cls"].save_lora_weights( - save_directory=self.args.output_dir, - transformer_lora_layers=transformer_lora_layers, - ) + if self.args.training_type == "lora": + transformer_lora_layers = get_peft_model_state_dict(transformer) + + self.model_config["pipeline_cls"].save_lora_weights( + save_directory=self.args.output_dir, + transformer_lora_layers=transformer_lora_layers, + ) + else: + transformer.save_pretrained(os.path.join(self.args.output_dir, "transformer")) self.validate(step=global_step, final_validation=True) @@ -910,35 +922,7 @@ class Trainer: memory_statistics = get_memory_statistics() logger.info(f"Memory before validation start: {json.dumps(memory_statistics, indent=4)}") - if not final_validation: - pipeline = self.model_config["initialize_pipeline"]( - model_id=self.args.pretrained_model_name_or_path, - tokenizer=self.tokenizer, - text_encoder=self.text_encoder, - tokenizer_2=self.tokenizer_2, - text_encoder_2=self.text_encoder_2, - transformer=unwrap_model(accelerator, self.transformer), - vae=self.vae, - device=accelerator.device, - revision=self.args.revision, - cache_dir=self.args.cache_dir, - enable_slicing=self.args.enable_slicing, - enable_tiling=self.args.enable_tiling, - enable_model_cpu_offload=self.args.enable_model_cpu_offload, - ) - else: - # `torch_dtype` is manually set within `initialize_pipeline()`. - self._delete_components() - pipeline = self.model_config["initialize_pipeline"]( - model_id=self.args.pretrained_model_name_or_path, - device=accelerator.device, - revision=self.args.revision, - cache_dir=self.args.cache_dir, - enable_slicing=self.args.enable_slicing, - enable_tiling=self.args.enable_tiling, - enable_model_cpu_offload=self.args.enable_model_cpu_offload, - ) - pipeline.load_lora_weights(self.args.output_dir) + pipeline = self._get_and_prepare_pipeline_for_validation(final_validation=final_validation) all_processes_artifacts = [] prompts_to_filenames = {} @@ -953,7 +937,7 @@ class Trainer: height = self.args.validation_heights[i] width = self.args.validation_widths[i] num_frames = self.args.validation_num_frames[i] - + frame_rate = self.args.validation_frame_rate if image is not None: image = load_image(image) if video is not None: @@ -971,6 +955,7 @@ class Trainer: height=height, width=width, num_frames=num_frames, + frame_rate=frame_rate, num_videos_per_prompt=self.args.num_validation_videos_per_prompt, generator=torch.Generator(device=accelerator.device).manual_seed( self.args.seed if self.args.seed is not None else 0 @@ -1010,7 +995,7 @@ class Trainer: elif artifact_type == "video": logger.debug(f"Saving video to {filename}") # TODO: this should be configurable here as well as in validation runs where we call the pipeline that has `fps`. - export_to_video(artifact_value, filename, fps=15) + export_to_video(artifact_value, filename, fps=frame_rate) artifact_value = wandb.Video(filename, caption=prompt) all_processes_artifacts.append(artifact_value) @@ -1144,3 +1129,56 @@ class Trainer: elif self.state.accelerator.mixed_precision == "bf16": weight_dtype = torch.bfloat16 return weight_dtype + + def _get_and_prepare_pipeline_for_validation(self, final_validation: bool = False) -> DiffusionPipeline: + accelerator = self.state.accelerator + if not final_validation: + pipeline = self.model_config["initialize_pipeline"]( + model_id=self.args.pretrained_model_name_or_path, + tokenizer=self.tokenizer, + text_encoder=self.text_encoder, + tokenizer_2=self.tokenizer_2, + text_encoder_2=self.text_encoder_2, + transformer=unwrap_model(accelerator, self.transformer), + vae=self.vae, + device=accelerator.device, + revision=self.args.revision, + cache_dir=self.args.cache_dir, + enable_slicing=self.args.enable_slicing, + enable_tiling=self.args.enable_tiling, + enable_model_cpu_offload=self.args.enable_model_cpu_offload, + ) + else: + self._delete_components() + + # Load the transformer weights from the final checkpoint if performing full-finetune + transformer = None + if self.args.training_type == "full-finetune": + transformer = self.model_config["load_diffusion_models"](model_id=self.args.output_dir)["transformer"] + + pipeline = self.model_config["initialize_pipeline"]( + model_id=self.args.pretrained_model_name_or_path, + transformer=transformer, + device=accelerator.device, + revision=self.args.revision, + cache_dir=self.args.cache_dir, + enable_slicing=self.args.enable_slicing, + enable_tiling=self.args.enable_tiling, + enable_model_cpu_offload=self.args.enable_model_cpu_offload, + ) + + # Load the LoRA weights if performing LoRA finetuning + if self.args.training_type == "lora": + pipeline.load_lora_weights(self.args.output_dir) + + return pipeline + + def _disable_grad_for_components(self, components: List[torch.nn.Module]): + for component in components: + if component is not None: + component.requires_grad_(False) + + def _enable_grad_for_components(self, components: List[torch.nn.Module]): + for component in components: + if component is not None: + component.requires_grad_(True)