mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
Suggestions Addressed
Renamed to SFT Added 2 other models. Testing required.
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user