Improve download

This commit is contained in:
litagin02
2024-05-25 19:28:22 +09:00
parent 5af4e2c6ac
commit a4f28e4f6f

View File

@@ -1,5 +1,7 @@
import shutil
import gradio as gr
from huggingface_hub import hf_hub_download, snapshot_download
from huggingface_hub import snapshot_download
from config import get_path_config
@@ -14,7 +16,7 @@ Hugging Face 🤗 に公開されているモデルをダウンロードして
例:
- `https://huggingface.co/username/my_sbv2_model`を指定すると、`model_assets/username-my_sbv2_model`に全体がダウンロードされます。
- `https://huggingface.co/username/my_sbv2_models/tree/main/model1`を指定すると、`model_assets/username-my_sbv2_models/model1`に`model1`フォルダのみがダウンロードされます。
- `https://huggingface.co/username/my_sbv2_models/tree/main/model1`を指定すると、`model_assets/username-my_sbv2_models/model1`に`model1`フォルダのみがダウンロードされますこの場合、model1フォルダ
**注意**
@@ -38,14 +40,24 @@ def download_model(url: str):
repo_folder = repo_folder[:-1]
if repo_folder == "":
model_name = repo_id.replace("/", "-")
result = snapshot_download(repo_id, local_dir=assets_root / model_name)
local_dir = assets_root / model_name
result = snapshot_download(repo_id, local_dir=local_dir)
else:
model_name = repo_id.replace("/", "-")
model_name = repo_id.replace("/", "-") + "-" + repo_folder.split("/")[-1]
local_dir = assets_root / model_name
result = snapshot_download(
repo_id,
local_dir=assets_root / model_name,
local_dir=local_dir,
allow_patterns=[repo_folder + "/*"],
)
# Move the downloaded folder to the correct path
for item in (assets_root / model_name / repo_folder).iterdir():
shutil.move(item, assets_root / model_name)
shutil.rmtree(assets_root / model_name / repo_folder.split("/")[0])
# Remove local_dir/.huggingface
hf_dir = local_dir / ".huggingface"
if hf_dir.exists():
shutil.rmtree(local_dir / ".huggingface")
return f"ダウンロード完了: {result}"