mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
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:
@@ -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
@@ -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`.
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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 +1,2 @@
|
||||
from .cogvideox_lora import COGVIDEOX_T2V_LORA_CONFIG
|
||||
from .full_finetune import COGVIDEOX_T2V_FULL_FINETUNE_CONFIG
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 +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 +1,2 @@
|
||||
from .full_finetune import LTX_VIDEO_T2V_FULL_FINETUNE_CONFIG
|
||||
from .ltx_video_lora import LTX_VIDEO_T2V_LORA_CONFIG
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
+110
-72
@@ -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,6 +391,7 @@ class Trainer:
|
||||
if self.args.gradient_checkpointing:
|
||||
self.transformer.enable_gradient_checkpointing()
|
||||
|
||||
if self.args.training_type == "lora":
|
||||
transformer_lora_config = LoraConfig(
|
||||
r=self.args.rank,
|
||||
lora_alpha=self.args.lora_alpha,
|
||||
@@ -413,6 +399,8 @@ class Trainer:
|
||||
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,6 +420,7 @@ class Trainer:
|
||||
type(unwrap_model(self.state.accelerator, self.transformer)),
|
||||
):
|
||||
model = unwrap_model(self.state.accelerator, 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()
|
||||
|
||||
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,18 +456,22 @@ 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(
|
||||
transformer_cls_ = unwrap_model(self.state.accelerator, self.transformer).__class__
|
||||
|
||||
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")
|
||||
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)
|
||||
@@ -483,9 +484,11 @@ class Trainer:
|
||||
# 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":
|
||||
if self.args.mixed_precision == "fp16" and self.args.training_type == "lora":
|
||||
# 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_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)
|
||||
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user