Suggestions Addressed

Renamed to SFT
Added 2 other models.
Testing required.
This commit is contained in:
CrossProduct
2025-01-07 22:59:09 +00:00
parent 4cd5a8e53b
commit cb9381b40e
7 changed files with 69 additions and 35 deletions
+1 -1
View File
@@ -455,7 +455,7 @@ def _add_training_arguments(parser: argparse.ArgumentParser) -> None:
"--training_type",
type=str,
default=None,
help="Type of training to perform. Choose between ['lora','full_finetune']",
help="Type of training to perform. Choose between ['lora','sft']",
)
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
parser.add_argument(
+15
View File
@@ -325,3 +325,18 @@ COGVIDEOX_T2V_LORA_CONFIG = {
"forward_pass": forward_pass,
"validation": validation,
}
COGVIDEOX_T2V_SFT_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,
}
@@ -358,3 +358,17 @@ HUNYUAN_VIDEO_T2V_LORA_CONFIG = {
"forward_pass": forward_pass,
"validation": validation,
}
HUNYUAN_VIDEO_T2V_SFT_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,
}
+1 -2
View File
@@ -322,7 +322,7 @@ LTX_VIDEO_T2V_LORA_CONFIG = {
"validation": validation,
}
LTX_VIDEO_T2V_FT_CONFIG = {
LTX_VIDEO_T2V_SFT_CONFIG = {
"pipeline_cls": LTXPipeline,
"load_condition_models": load_condition_models,
"load_latent_models": load_latent_models,
@@ -335,4 +335,3 @@ LTX_VIDEO_T2V_FT_CONFIG = {
"forward_pass": forward_pass,
"validation": validation,
}
+6 -4
View File
@@ -1,20 +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, LTX_VIDEO_T2V_FT_CONFIG
from .cogvideox import COGVIDEOX_T2V_LORA_CONFIG, COGVIDEOX_T2V_SFT_CONFIG
from .hunyuan_video import HUNYUAN_VIDEO_T2V_LORA_CONFIG, HUNYUAN_VIDEO_T2V_SFT_CONFIG
from .ltx_video import LTX_VIDEO_T2V_LORA_CONFIG, LTX_VIDEO_T2V_SFT_CONFIG
SUPPORTED_MODEL_CONFIGS = {
"hunyuan_video": {
"lora": HUNYUAN_VIDEO_T2V_LORA_CONFIG,
"sft": HUNYUAN_VIDEO_T2V_SFT_CONFIG,
},
"ltx_video": {
"lora": LTX_VIDEO_T2V_LORA_CONFIG,
"finetune": LTX_VIDEO_T2V_FT_CONFIG,
"sft": LTX_VIDEO_T2V_SFT_CONFIG,
},
"cogvideox": {
"lora": COGVIDEOX_T2V_LORA_CONFIG,
"sft": COGVIDEOX_T2V_SFT_CONFIG,
},
}
+32 -26
View File
@@ -101,6 +101,7 @@ class Trainer:
# Components list
self.components = []
def prepare_dataset(self) -> None:
# TODO(aryan): Make a background process for fetching
logger.info("Initializing dataset and dataloader")
@@ -155,16 +156,17 @@ class Trainer:
self.transformer_config = self.transformer.config if self.transformer is not None else self.transformer_config
self.vae_config = self.vae.config if self.vae is not None else self.vae_config
self.components = [self.tokenizer,
self.tokenizer_2,
self.tokenizer_3,
self.text_encoder,
self.text_encoder_2,
self.text_encoder_3,
self.transformer,
self.unet,
self.vae]
self.components = [
self.tokenizer,
self.tokenizer_2,
self.tokenizer_3,
self.text_encoder,
self.text_encoder_2,
self.text_encoder_3,
self.transformer,
self.unet,
self.vae,
]
def _delete_components(self) -> None:
self.tokenizer = None
@@ -204,12 +206,12 @@ class Trainer:
if self.args.enable_tiling:
self.vae.enable_tiling()
def _disable_grad_for_components(self, components:list):
def _disable_grad_for_components(self, components: list):
for component in components:
if component is not None:
component.requires_grad_(False)
def _enable_grad_for_components(self, components:list):
def _enable_grad_for_components(self, components: list):
for component in components:
if component is not None:
component.requires_grad_(True)
@@ -262,11 +264,13 @@ class Trainer:
self._set_components(condition_components)
self._move_components_to_device()
self._disable_grad_for_components(components=[
self.text_encoder,
self.text_encoder_2,
self.text_encoder_3,
])
self._disable_grad_for_components(
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(
"Caption dropout is not supported with precomputation yet. This will be supported in the future."
@@ -378,20 +382,22 @@ class Trainer:
diffusion_components = self.model_config["load_diffusion_models"](**self._get_load_components_kwargs())
self._set_components(diffusion_components)
self._disable_grad_for_components(components=[
self.text_encoder,
self.text_encoder_2,
self.text_encoder_3,
self.vae,
])
if self.args.training_type == "full_finetune":
self._disable_grad_for_components(
components=[
self.text_encoder,
self.text_encoder_2,
self.text_encoder_3,
self.vae,
]
)
if self.args.training_type == "sft":
logger.info("Full Fine Tuning Enabled")
self._enable_grad_for_components(components=[self.transformer])
else:
logger.info("Lora Fine Tuning Enabled")
self._disable_grad_for_components(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)
-2
View File
@@ -4,11 +4,9 @@ import traceback
from finetrainers import Trainer, parse_arguments
from finetrainers.constants import FINETRAINERS_LOG_LEVEL
logger = logging.getLogger("finetrainers")
logger.setLevel(FINETRAINERS_LOG_LEVEL)
def main():
try:
import multiprocessing