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
- 🔥 **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
<div align="center">
| **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<sup>*</sup> |
| **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 | 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 |
</div>
<sub><sup>*</sup>Noted for the 5B variant.</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.
<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>
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
* 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)
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`.
+29
View File
@@ -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:
+8
View File
@@ -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:
+28 -1
View File
@@ -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:
+28 -2
View File
@@ -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:
+1
View File
@@ -1 +1,2 @@
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
# TODO(aryan): refactor into model specs for better re-use
COGVIDEOX_T2V_LORA_CONFIG = {
"pipeline_cls": CogVideoXPipeline,
"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
@@ -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}
# TODO(aryan): refactor into model specs for better re-use
HUNYUAN_VIDEO_T2V_LORA_CONFIG = {
"pipeline_cls": HunyuanVideoPipeline,
"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
+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,
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,
+6 -3
View File
@@ -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,
},
}
+145 -107
View File
@@ -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)