Full Finetuning for LTX possibily extended to other models. (#192)

* Full Finetuning for LTX possibily extended to other models.

* Change name of the flag

* Used disable grad for component on lora fine tuning enabled

* Suggestions Addressed
Renamed to SFT
Added 2 other models.
Testing required.

* Switching to Full FineTuning

* Run linter.

* parse subfolder when needed.

* tackle saving and loading hooks.

* tackle validation.

* fix subfolder bug.

* remove __class__.

* refactor

* remove unnecessary changes

* handle saving of final model weights correctly

* remove unnecessary changes

* LTX uses a default frame rate of 24 FPS
We need to modify the output validation framerate to match that value.
Add Framerate args.
Add Update video output and inference frame rate

* There was a results_args mapping that needed to be modified.

* update

* update README

* Update README.md

* update docs

* add training configuration in cogvideox

---------

Co-authored-by: Sayak Paul <spsayakpaul@gmail.com>
Co-authored-by: Aryan <aryan@huggingface.co>
Co-authored-by: Aryan <contact.aryanvs@gmail.com>
This commit is contained in:
ArEnSc
2025-01-13 13:25:39 -05:00
committed by GitHub
parent a3898075d1
commit e5df80cc36
17 changed files with 365 additions and 131 deletions
+9 -9
View File
@@ -12,7 +12,8 @@ FineTrainers is a work-in-progress library to support (accessible) training of v
## News ## 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-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! - 🔥 **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
<div align="center"> <div align="center">
| **Model Name** | **Tasks** | **Min. GPU VRAM** | | **Model Name** | **Tasks** | **Min. LoRA VRAM<sup>*</sup>** | **Min. Full Finetuning VRAM<sup>^</sup>** |
|:---:|:---:|:---:| |:------------------------------------------------:|:-------------:|:----------------------------------:|:---------------------------------------------:|
| [LTX-Video](./docs/training/ltx_video.md) | Text-to-Video | 11 GB | | [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 | | [HunyuanVideo](./docs/training/hunyuan_video.md) | Text-to-Video | 42 GB | OOM |
| [CogVideoX](./docs/training/cogvideox.md) | Text-to-Video | 12GB<sup>*</sup> | | [CogVideoX-5b](./docs/training/cogvideox.md) | Text-to-Video | 21 GB | 53 GB |
</div> </div>
<sub><sup>*</sup>Noted for the 5B variant.</sub> <sub><sup>*</sup>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).</sub><br/>
<sub><sup>^</sup>Noted for training-only, no validation, at resolution `49x512x768`, with pre-computation, using bf16 weights & gradient checkpointing.</sub>
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.
If you would like to use a custom dataset, refer to the dataset preparation guide [here](./docs/dataset/README.md). If you would like to use a custom dataset, refer to the dataset preparation guide [here](./docs/dataset/README.md).
+14 -8
View File
@@ -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 This directory contains the training-related specifications for all the models we support in `finetrainers`. Each model page has:
* inference example - an example training command
* numbers on memory consumption - 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: 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 + --validation_steps 100
``` ```
## Model-specific docs Supported models:
- [CogVideoX](./cogvideox.md)
- [LTX-Video](./ltx_video.md)
- [HunyuanVideo](./hunyuan_video.md)
* [CogVideoX](./cogvideox.md) Supported training types:
* [LTX-Video](./ltx_video.md) - LoRA (`--training_type lora`)
* [HunyuanVideo](./hunyuan_video.md) - Full finetuning (`--training_type full-finetune`)
Arguments for training are well-documented in the code. For more information, please run `python train.py --help`.
+29
View File
@@ -2,6 +2,8 @@
## Training ## Training
For LoRA training, specify `--training_type lora`. For full finetuning, specify `--training_type full-finetune`.
```bash ```bash
#!/bin/bash #!/bin/bash
export WANDB_MODE="offline" export WANDB_MODE="offline"
@@ -84,6 +86,8 @@ echo -ne "-------------------- Finished executing script --------------------\n\
## Memory Usage ## Memory Usage
### LoRA
LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x480x720` resolutions, **with precomputation**: 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 validation end | 11.145 | 28.324 |
| after training end | 11.144 | 11.592 | | 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 ## 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: 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:
+8
View File
@@ -2,6 +2,8 @@
## Training ## Training
For LoRA training, specify `--training_type lora`. For full finetuning, specify `--training_type full-finetune`.
```bash ```bash
#!/bin/bash #!/bin/bash
@@ -87,6 +89,8 @@ echo -ne "-------------------- Finished executing script --------------------\n\
## Memory Usage ## Memory Usage
### LoRA
LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolutions, **without precomputation**: 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. 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 ## 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: 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:
+28 -1
View File
@@ -2,7 +2,7 @@
## Training ## Training
Provided you have a dataset: For LoRA training, specify `--training_type lora`. For full finetuning, specify `--training_type full-finetune`.
```bash ```bash
#!/bin/bash #!/bin/bash
@@ -88,6 +88,8 @@ echo -ne "-------------------- Finished executing script --------------------\n\
## Memory Usage ## Memory Usage
### LoRA
LoRA with rank 128, batch size 1, gradient checkpointing, optimizer adamw, `49x512x768` resolution, **without precomputation**: 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. 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 ## 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: 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:
+28 -2
View File
@@ -207,6 +207,8 @@ class Args:
Perform validation every `n` training steps. Perform validation every `n` training steps.
enable_model_cpu_offload (`bool`, defaults to `False`): enable_model_cpu_offload (`bool`, defaults to `False`):
Whether or not to offload different modeling components to CPU during validation. 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 MISCELLANEOUS ARGUMENTS
----------------------- -----------------------
@@ -319,6 +321,7 @@ class Args:
validation_every_n_epochs: Optional[int] = None validation_every_n_epochs: Optional[int] = None
validation_every_n_steps: Optional[int] = None validation_every_n_steps: Optional[int] = None
enable_model_cpu_offload: bool = False enable_model_cpu_offload: bool = False
validation_frame_rate: int = 25
# Miscellaneous arguments # Miscellaneous arguments
tracker_name: str = "finetrainers" tracker_name: str = "finetrainers"
@@ -417,6 +420,7 @@ class Args:
"validation_every_n_epochs": self.validation_every_n_epochs, "validation_every_n_epochs": self.validation_every_n_epochs,
"validation_every_n_steps": self.validation_every_n_steps, "validation_every_n_steps": self.validation_every_n_steps,
"enable_model_cpu_offload": self.enable_model_cpu_offload, "enable_model_cpu_offload": self.enable_model_cpu_offload,
"validation_frame_rate": self.validation_frame_rate,
}, },
"miscellaneous_arguments": { "miscellaneous_arguments": {
"tracker_name": self.tracker_name, "tracker_name": self.tracker_name,
@@ -460,6 +464,7 @@ def parse_arguments() -> Args:
def validate_args(args: Args): def validate_args(args: Args):
_validate_training_args(args)
_validate_validation_args(args) _validate_validation_args(args)
@@ -678,8 +683,9 @@ def _add_training_arguments(parser: argparse.ArgumentParser) -> None:
parser.add_argument( parser.add_argument(
"--training_type", "--training_type",
type=str, type=str,
choices=["lora", "full-finetune"],
required=True, 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("--seed", type=int, default=None, help="A seed for reproducible training.")
parser.add_argument( 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.", help="The lora_alpha to compute scaling factor (lora_alpha / rank) for LoRA matrices.",
) )
parser.add_argument( 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( parser.add_argument(
"--gradient_accumulation_steps", "--gradient_accumulation_steps",
@@ -890,6 +900,12 @@ def _add_validation_arguments(parser: argparse.ArgumentParser) -> None:
default=None, default=None,
help="Run validation every X training steps. Validation consists of running the validation prompt `args.num_validation_videos` times.", 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( parser.add_argument(
"--enable_model_cpu_offload", "--enable_model_cpu_offload",
action="store_true", 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_epochs = args.validation_epochs
result_args.validation_every_n_steps = args.validation_steps result_args.validation_every_n_steps = args.validation_steps
result_args.enable_model_cpu_offload = args.enable_model_cpu_offload result_args.enable_model_cpu_offload = args.enable_model_cpu_offload
result_args.validation_frame_rate = args.validation_frame_rate
# Miscellaneous arguments # Miscellaneous arguments
result_args.tracker_name = args.tracker_name result_args.tracker_name = args.tracker_name
@@ -1100,6 +1117,15 @@ def _map_to_args_type(args: Dict[str, Any]) -> Args:
return result_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): def _validate_validation_args(args: Args):
assert args.validation_prompts is not None, "Validation prompts are required for validation" assert args.validation_prompts is not None, "Validation prompts are required for validation"
if args.validation_images is not None: if args.validation_images is not None:
+1
View File
@@ -1 +1,2 @@
from .cogvideox_lora import COGVIDEOX_T2V_LORA_CONFIG from .cogvideox_lora import COGVIDEOX_T2V_LORA_CONFIG
from .full_finetune import COGVIDEOX_T2V_FULL_FINETUNE_CONFIG
+1
View File
@@ -311,6 +311,7 @@ def _pad_frames(latents: torch.Tensor, patch_size_t: int):
return latents return latents
# TODO(aryan): refactor into model specs for better re-use
COGVIDEOX_T2V_LORA_CONFIG = { COGVIDEOX_T2V_LORA_CONFIG = {
"pipeline_cls": CogVideoXPipeline, "pipeline_cls": CogVideoXPipeline,
"load_condition_models": load_condition_models, "load_condition_models": load_condition_models,
+32
View File
@@ -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,
}
+1
View File
@@ -1 +1,2 @@
from .full_finetune import HUNYUAN_VIDEO_T2V_FULL_FINETUNE_CONFIG
from .hunyuan_video_lora import HUNYUAN_VIDEO_T2V_LORA_CONFIG from .hunyuan_video_lora import HUNYUAN_VIDEO_T2V_LORA_CONFIG
@@ -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,
}
@@ -345,6 +345,7 @@ def _get_clip_prompt_embeds(
return {"pooled_prompt_embeds": prompt_embeds} return {"pooled_prompt_embeds": prompt_embeds}
# TODO(aryan): refactor into model specs for better re-use
HUNYUAN_VIDEO_T2V_LORA_CONFIG = { HUNYUAN_VIDEO_T2V_LORA_CONFIG = {
"pipeline_cls": HunyuanVideoPipeline, "pipeline_cls": HunyuanVideoPipeline,
"load_condition_models": load_condition_models, "load_condition_models": load_condition_models,
+1
View File
@@ -1 +1,2 @@
from .full_finetune import LTX_VIDEO_T2V_FULL_FINETUNE_CONFIG
from .ltx_video_lora import LTX_VIDEO_T2V_LORA_CONFIG from .ltx_video_lora import LTX_VIDEO_T2V_LORA_CONFIG
+30
View File
@@ -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,
}
+1 -1
View File
@@ -225,7 +225,7 @@ def validation(
height: Optional[int] = None, height: Optional[int] = None,
width: Optional[int] = None, width: Optional[int] = None,
num_frames: Optional[int] = None, num_frames: Optional[int] = None,
frame_rate: int = 25, frame_rate: int = 24,
num_videos_per_prompt: int = 1, num_videos_per_prompt: int = 1,
generator: Optional[torch.Generator] = None, generator: Optional[torch.Generator] = None,
**kwargs, **kwargs,
+6 -3
View File
@@ -1,19 +1,22 @@
from typing import Any, Dict from typing import Any, Dict
from .cogvideox import COGVIDEOX_T2V_LORA_CONFIG from .cogvideox import COGVIDEOX_T2V_FULL_FINETUNE_CONFIG, COGVIDEOX_T2V_LORA_CONFIG
from .hunyuan_video import HUNYUAN_VIDEO_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_LORA_CONFIG from .ltx_video import LTX_VIDEO_T2V_FULL_FINETUNE_CONFIG, LTX_VIDEO_T2V_LORA_CONFIG
SUPPORTED_MODEL_CONFIGS = { SUPPORTED_MODEL_CONFIGS = {
"hunyuan_video": { "hunyuan_video": {
"lora": HUNYUAN_VIDEO_T2V_LORA_CONFIG, "lora": HUNYUAN_VIDEO_T2V_LORA_CONFIG,
"full-finetune": HUNYUAN_VIDEO_T2V_FULL_FINETUNE_CONFIG,
}, },
"ltx_video": { "ltx_video": {
"lora": LTX_VIDEO_T2V_LORA_CONFIG, "lora": LTX_VIDEO_T2V_LORA_CONFIG,
"full-finetune": LTX_VIDEO_T2V_FULL_FINETUNE_CONFIG,
}, },
"cogvideox": { "cogvideox": {
"lora": COGVIDEOX_T2V_LORA_CONFIG, "lora": COGVIDEOX_T2V_LORA_CONFIG,
"full-finetune": COGVIDEOX_T2V_FULL_FINETUNE_CONFIG,
}, },
} }
+145 -107
View File
@@ -5,7 +5,7 @@ import os
import random import random
from datetime import datetime, timedelta from datetime import datetime, timedelta
from pathlib import Path from pathlib import Path
from typing import Any, Dict from typing import Any, Dict, List
import diffusers import diffusers
import torch import torch
@@ -21,6 +21,7 @@ from accelerate.utils import (
gather_object, gather_object,
set_seed, set_seed,
) )
from diffusers import DiffusionPipeline
from diffusers.configuration_utils import FrozenDict from diffusers.configuration_utils import FrozenDict
from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution
from diffusers.optimization import get_scheduler 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()) condition_components = self.model_config["load_condition_models"](**self._get_load_components_kwargs())
self._set_components(condition_components) self._set_components(condition_components)
self._move_components_to_device() self._move_components_to_device()
self._disable_grad_for_components([self.text_encoder, self.text_encoder_2, self.text_encoder_3])
# 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)
if self.args.caption_dropout_p > 0 and self.args.caption_dropout_technique == "empty": if self.args.caption_dropout_p > 0 and self.args.caption_dropout_technique == "empty":
logger.warning( logger.warning(
@@ -305,12 +297,7 @@ class Trainer:
latent_components = self.model_config["load_latent_models"](**self._get_load_components_kwargs()) latent_components = self.model_config["load_latent_models"](**self._get_load_components_kwargs())
self._set_components(latent_components) self._set_components(latent_components)
self._move_components_to_device() self._move_components_to_device()
self._disable_grad_for_components([self.vae])
# 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)
if self.vae is not None: if self.vae is not None:
if self.args.enable_slicing: if self.args.enable_slicing:
@@ -371,24 +358,22 @@ class Trainer:
diffusion_components = self.model_config["load_diffusion_models"](**self._get_load_components_kwargs()) diffusion_components = self.model_config["load_diffusion_models"](**self._get_load_components_kwargs())
self._set_components(diffusion_components) self._set_components(diffusion_components)
# TODO(aryan): refactor later. for now only lora is supported components = [self.text_encoder, self.text_encoder_2, self.text_encoder_3, self.vae]
components_to_disable_grads = [ self._disable_grad_for_components(components)
self.text_encoder,
self.text_encoder_2, if self.args.training_type == "full-finetune":
self.text_encoder_3, logger.info("Finetuning transformer with no additional parameters")
self.transformer, self._enable_grad_for_components([self.transformer])
self.vae, else:
] logger.info("Finetuning transformer with PEFT parameters")
for component in components_to_disable_grads: self._disable_grad_for_components([self.transformer])
if component is not None:
component.requires_grad_(False)
# For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision # 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. # 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) weight_dtype = self._get_training_dtype(accelerator=self.state.accelerator)
if torch.backends.mps.is_available() and 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. # Due to pytorch#99272, MPS does not yet support bfloat16.
raise ValueError( raise ValueError(
"Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." "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: if self.args.gradient_checkpointing:
self.transformer.enable_gradient_checkpointing() self.transformer.enable_gradient_checkpointing()
transformer_lora_config = LoraConfig( if self.args.training_type == "lora":
r=self.args.rank, transformer_lora_config = LoraConfig(
lora_alpha=self.args.lora_alpha, r=self.args.rank,
init_lora_weights=True, lora_alpha=self.args.lora_alpha,
target_modules=self.args.target_modules, init_lora_weights=True,
) target_modules=self.args.target_modules,
self.transformer.add_adapter(transformer_lora_config) )
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 # 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(): if self.args.allow_tf32 and torch.cuda.is_available():
@@ -432,7 +420,8 @@ class Trainer:
type(unwrap_model(self.state.accelerator, self.transformer)), type(unwrap_model(self.state.accelerator, self.transformer)),
): ):
model = unwrap_model(self.state.accelerator, model) 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: else:
raise ValueError(f"Unexpected save model: {model.__class__}") raise ValueError(f"Unexpected save model: {model.__class__}")
@@ -440,10 +429,18 @@ class Trainer:
if weights: if weights:
weights.pop() weights.pop()
self.model_config["pipeline_cls"].save_lora_weights( if self.args.training_type == "lora":
output_dir, self.model_config["pipeline_cls"].save_lora_weights(
transformer_lora_layers=transformer_lora_layers_to_save, 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): def load_model_hook(models, input_dir):
if not self.state.accelerator.distributed_type == DistributedType.DEEPSPEED: 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__}" f"Unexpected save model: {unwrap_model(self.state.accelerator, model).__class__}"
) )
else: else:
transformer_ = unwrap_model(self.state.accelerator, self.transformer).__class__.from_pretrained( transformer_cls_ = unwrap_model(self.state.accelerator, self.transformer).__class__
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) if self.args.training_type == "lora":
transformer_state_dict = { transformer_ = transformer_cls_.from_pretrained(
f'{k.replace("transformer.", "")}': v self.args.pretrained_model_name_or_path, subfolder="transformer"
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}. "
) )
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 # Make sure the trainable params are in float32. This is again needed since the base models
# are in `weight_dtype`. More details: # are in `weight_dtype`. More details:
# https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804 # https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804
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 # only upcast trainable parameters (LoRA) into fp32
cast_training_params([transformer_]) 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_save_state_pre_hook(save_model_hook)
self.state.accelerator.register_load_state_pre_hook(load_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 self.state.train_steps = self.args.train_steps
# Make sure the trainable params are in float32 # 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 # only upcast trainable parameters (LoRA) into fp32
cast_training_params([self.transformer], dtype=torch.float32) cast_training_params([self.transformer], dtype=torch.float32)
@@ -510,13 +513,13 @@ class Trainer:
* self.state.accelerator.num_processes * 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 = { transformer_parameters_with_lr = {
"params": transformer_lora_parameters, "params": transformer_trainable_parameters,
"lr": self.state.learning_rate, "lr": self.state.learning_rate,
} }
params_to_optimize = [transformer_parameters_with_lr] 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 = ( use_deepspeed_opt = (
self.state.accelerator.state.deepspeed_plugin is not None self.state.accelerator.state.deepspeed_plugin is not None
@@ -608,6 +611,12 @@ class Trainer:
) )
self.vae_config = FrozenDict(**vae_config) 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.state.train_batch_size = (
self.args.batch_size * self.state.accelerator.num_processes * self.args.gradient_accumulation_steps self.args.batch_size * self.state.accelerator.num_processes * self.args.gradient_accumulation_steps
) )
@@ -872,14 +881,17 @@ class Trainer:
accelerator.wait_for_everyone() accelerator.wait_for_everyone()
if accelerator.is_main_process: if accelerator.is_main_process:
# TODO: consider factoring this out when supporting other types of training algos. transformer = unwrap_model(accelerator, self.transformer)
self.transformer = unwrap_model(accelerator, self.transformer)
transformer_lora_layers = get_peft_model_state_dict(self.transformer)
self.model_config["pipeline_cls"].save_lora_weights( if self.args.training_type == "lora":
save_directory=self.args.output_dir, transformer_lora_layers = get_peft_model_state_dict(transformer)
transformer_lora_layers=transformer_lora_layers,
) 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) self.validate(step=global_step, final_validation=True)
@@ -910,35 +922,7 @@ class Trainer:
memory_statistics = get_memory_statistics() memory_statistics = get_memory_statistics()
logger.info(f"Memory before validation start: {json.dumps(memory_statistics, indent=4)}") logger.info(f"Memory before validation start: {json.dumps(memory_statistics, indent=4)}")
if not final_validation: pipeline = self._get_and_prepare_pipeline_for_validation(final_validation=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)
all_processes_artifacts = [] all_processes_artifacts = []
prompts_to_filenames = {} prompts_to_filenames = {}
@@ -953,7 +937,7 @@ class Trainer:
height = self.args.validation_heights[i] height = self.args.validation_heights[i]
width = self.args.validation_widths[i] width = self.args.validation_widths[i]
num_frames = self.args.validation_num_frames[i] num_frames = self.args.validation_num_frames[i]
frame_rate = self.args.validation_frame_rate
if image is not None: if image is not None:
image = load_image(image) image = load_image(image)
if video is not None: if video is not None:
@@ -971,6 +955,7 @@ class Trainer:
height=height, height=height,
width=width, width=width,
num_frames=num_frames, num_frames=num_frames,
frame_rate=frame_rate,
num_videos_per_prompt=self.args.num_validation_videos_per_prompt, num_videos_per_prompt=self.args.num_validation_videos_per_prompt,
generator=torch.Generator(device=accelerator.device).manual_seed( generator=torch.Generator(device=accelerator.device).manual_seed(
self.args.seed if self.args.seed is not None else 0 self.args.seed if self.args.seed is not None else 0
@@ -1010,7 +995,7 @@ class Trainer:
elif artifact_type == "video": elif artifact_type == "video":
logger.debug(f"Saving video to {filename}") 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`. # 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) artifact_value = wandb.Video(filename, caption=prompt)
all_processes_artifacts.append(artifact_value) all_processes_artifacts.append(artifact_value)
@@ -1144,3 +1129,56 @@ class Trainer:
elif self.state.accelerator.mixed_precision == "bf16": elif self.state.accelerator.mixed_precision == "bf16":
weight_dtype = torch.bfloat16 weight_dtype = torch.bfloat16
return weight_dtype 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)