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