diff --git a/finetrainers/args.py b/finetrainers/args.py index 224b9f1..bcd0076 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -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( diff --git a/finetrainers/cogvideox/cogvideox_lora.py b/finetrainers/cogvideox/cogvideox_lora.py index c3b754a..a536cab 100644 --- a/finetrainers/cogvideox/cogvideox_lora.py +++ b/finetrainers/cogvideox/cogvideox_lora.py @@ -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, +} diff --git a/finetrainers/hunyuan_video/hunyuan_video_lora.py b/finetrainers/hunyuan_video/hunyuan_video_lora.py index 9bfea53..34ce576 100644 --- a/finetrainers/hunyuan_video/hunyuan_video_lora.py +++ b/finetrainers/hunyuan_video/hunyuan_video_lora.py @@ -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, +} diff --git a/finetrainers/ltx_video/ltx_video_lora.py b/finetrainers/ltx_video/ltx_video_lora.py index 4b14533..77ec50e 100644 --- a/finetrainers/ltx_video/ltx_video_lora.py +++ b/finetrainers/ltx_video/ltx_video_lora.py @@ -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, } - diff --git a/finetrainers/models.py b/finetrainers/models.py index dab4de3..e753c4a 100644 --- a/finetrainers/models.py +++ b/finetrainers/models.py @@ -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, }, } diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index e948411..bfcbaeb 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -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) diff --git a/train.py b/train.py index 088c061..32fe903 100644 --- a/train.py +++ b/train.py @@ -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