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>
26 lines
869 B
Python
26 lines
869 B
Python
import importlib
|
|
import json
|
|
import os
|
|
|
|
from huggingface_hub import hf_hub_download
|
|
|
|
|
|
def resolve_vae_cls_from_ckpt_path(ckpt_path, **kwargs):
|
|
ckpt_path = str(ckpt_path)
|
|
if os.path.exists(str(ckpt_path)) and os.path.isdir(ckpt_path):
|
|
index_path = os.path.join(ckpt_path, "model_index.json")
|
|
else:
|
|
revision = kwargs.get("revision", None)
|
|
cache_dir = kwargs.get("cache_dir", None)
|
|
index_path = hf_hub_download(
|
|
repo_id=ckpt_path, filename="model_index.json", revision=revision, cache_dir=cache_dir
|
|
)
|
|
|
|
with open(index_path, "r") as f:
|
|
model_index_dict = json.load(f)
|
|
assert "vae" in model_index_dict, "No VAE found in the modelx index dict."
|
|
|
|
vae_cls_config = model_index_dict["vae"]
|
|
library = importlib.import_module(vae_cls_config[0])
|
|
return getattr(library, vae_cls_config[1])
|