mindfull / scripts /download_models.py
IamSamk
Mindfull Gradio Space deploy
27caffe
Raw
History Blame Contribute Delete
2.12 kB
#!/usr/bin/env python3
"""Download ML weights from Hugging Face Hub into local paths expected by the app."""
from __future__ import annotations
import os
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
# Update these after you upload models to your Hugging Face account.
DEFAULT_REPOS = {
"f5tts": "IamSamkk/police-bot-f5tts",
"sadtalker": "IamSamkk/sadtalker-checkpoints",
"gfpgan": "IamSamkk/gfpgan-weights",
}
LOCAL_DIRS = {
"f5tts": ROOT / "my_finetuned_model",
"sadtalker": ROOT / "sadtalker+wav2lip" / "sadtalker" / "checkpoints",
"gfpgan": ROOT / "gfpgan" / "weights",
}
ENV_REPO_KEYS = {
"f5tts": "HF_REPO_F5TTS",
"sadtalker": "HF_REPO_SADTALKER",
"gfpgan": "HF_REPO_GFPGAN",
}
def _repo_id(key: str) -> str:
env_key = ENV_REPO_KEYS[key]
return os.environ.get(env_key, DEFAULT_REPOS[key])
def download_all(token: str | None = None, only: list[str] | None = None) -> None:
try:
from huggingface_hub import snapshot_download
except ImportError:
print("Install huggingface_hub: pip install huggingface_hub", file=sys.stderr)
sys.exit(1)
keys = only if only else list(DEFAULT_REPOS.keys())
for key in keys:
if key not in DEFAULT_REPOS:
print(f"Unknown model key: {key}", file=sys.stderr)
sys.exit(1)
repo_id = _repo_id(key)
target = LOCAL_DIRS[key]
target.mkdir(parents=True, exist_ok=True)
print(f"Downloading {repo_id} -> {target}")
snapshot_download(
repo_id=repo_id,
local_dir=str(target),
token=token or os.environ.get("HF_TOKEN"),
)
print("Done.")
def main() -> None:
import argparse
parser = argparse.ArgumentParser(description="Download model weights from Hugging Face Hub")
parser.add_argument(
"--only",
nargs="+",
choices=list(DEFAULT_REPOS.keys()),
help="Download only selected model groups",
)
args = parser.parse_args()
download_all(only=args.only)
if __name__ == "__main__":
main()