mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
e5df80cc36
* 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>
34 lines
1.3 KiB
Python
34 lines
1.3 KiB
Python
from typing import Any, Dict
|
|
|
|
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,
|
|
},
|
|
}
|
|
|
|
|
|
def get_config_from_model_name(model_name: str, training_type: str) -> Dict[str, Any]:
|
|
if model_name not in SUPPORTED_MODEL_CONFIGS:
|
|
raise ValueError(
|
|
f"Model {model_name} not supported. Supported models are: {list(SUPPORTED_MODEL_CONFIGS.keys())}"
|
|
)
|
|
if training_type not in SUPPORTED_MODEL_CONFIGS[model_name]:
|
|
raise ValueError(
|
|
f"Training type {training_type} not supported for model {model_name}. Supported training types are: {list(SUPPORTED_MODEL_CONFIGS[model_name].keys())}"
|
|
)
|
|
return SUPPORTED_MODEL_CONFIGS[model_name][training_type]
|