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)