mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
b8352abf70
* support cog t2v. * generator. * updates * style * fixes * fix padding frames. Co-authored-by: zRzRzRzRzRzRzR <Yuxuan.Zhang2104@student.xjtlu.edu.cn> * revert changes related to generator. * refactor a lot of things. * accept revision and cache_dir. * remove unused var * refactoring fixes. * refactor * update * update --------- Co-authored-by: zRzRzRzRzRzRzR <Yuxuan.Zhang2104@student.xjtlu.edu.cn> Co-authored-by: Aryan <aryan@huggingface.co>
31 lines
1.0 KiB
Python
31 lines
1.0 KiB
Python
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
|
|
|
|
|
|
SUPPORTED_MODEL_CONFIGS = {
|
|
"hunyuan_video": {
|
|
"lora": HUNYUAN_VIDEO_T2V_LORA_CONFIG,
|
|
},
|
|
"ltx_video": {
|
|
"lora": LTX_VIDEO_T2V_LORA_CONFIG,
|
|
},
|
|
"cogvideox": {
|
|
"lora": COGVIDEOX_T2V_LORA_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]
|