hugging-apps/echo-memory
0
1from huggingface_hub import hf_hub_download2from modelscope import snapshot_download3import os, shutil4from typing_extensions import Literal, TypeAlias5from typing import List6from ..configs.model_config import preset_models_on_huggingface, preset_models_on_modelscope, Preset_model_id7 8 9def download_from_modelscope(model_id, origin_file_path, local_dir):10 os.makedirs(local_dir, exist_ok=True)11 file_name = os.path.basename(origin_file_path)12 if file_name in os.listdir(local_dir):13 print(f" {file_name} has been already in {local_dir}.")14 else:15 print(f" Start downloading {os.path.join(local_dir, file_name)}")16 snapshot_download(model_id, allow_file_pattern=origin_file_path, local_dir=local_dir)17 downloaded_file_path = os.path.join(local_dir, origin_file_path)18 target_file_path = os.path.join(local_dir, os.path.split(origin_file_path)[-1])19 if downloaded_file_path != target_file_path:20 shutil.move(downloaded_file_path, target_file_path)21 shutil.rmtree(os.path.join(local_dir, origin_file_path.split("/")[0]))22 23 24def download_from_huggingface(model_id, origin_file_path, local_dir):25 os.makedirs(local_dir, exist_ok=True)26 file_name = os.path.basename(origin_file_path)27 if file_name in os.listdir(local_dir):28 print(f" {file_name} has been already in {local_dir}.")29 else:30 print(f" Start downloading {os.path.join(local_dir, file_name)}")31 hf_hub_download(model_id, origin_file_path, local_dir=local_dir)32 downloaded_file_path = os.path.join(local_dir, origin_file_path)33 target_file_path = os.path.join(local_dir, file_name)34 if downloaded_file_path != target_file_path:35 shutil.move(downloaded_file_path, target_file_path)36 shutil.rmtree(os.path.join(local_dir, origin_file_path.split("/")[0]))37 38 39Preset_model_website: TypeAlias = Literal[40 "HuggingFace",41 "ModelScope",42]43website_to_preset_models = {44 "HuggingFace": preset_models_on_huggingface,45 "ModelScope": preset_models_on_modelscope,46}47website_to_download_fn = {48 "HuggingFace": download_from_huggingface,49 "ModelScope": download_from_modelscope,50}51 52 53def download_customized_models(54 model_id,55 origin_file_path,56 local_dir,57 downloading_priority: List[Preset_model_website] = ["ModelScope", "HuggingFace"],58):59 downloaded_files = []60 for website in downloading_priority:61 # Check if the file is downloaded.62 file_to_download = os.path.join(local_dir, os.path.basename(origin_file_path))63 if file_to_download in downloaded_files:64 continue65 # Download66 website_to_download_fn[website](model_id, origin_file_path, local_dir)67 if os.path.basename(origin_file_path) in os.listdir(local_dir):68 downloaded_files.append(file_to_download)69 return downloaded_files70 71 72def download_models(73 model_id_list: List[Preset_model_id] = [],74 downloading_priority: List[Preset_model_website] = ["ModelScope", "HuggingFace"],75):76 print(f"Downloading models: {model_id_list}")77 downloaded_files = []78 load_files = []79 80 for model_id in model_id_list:81 for website in downloading_priority:82 if model_id in website_to_preset_models[website]:83 84 # Parse model metadata85 model_metadata = website_to_preset_models[website][model_id]86 if isinstance(model_metadata, list):87 file_data = model_metadata88 else:89 file_data = model_metadata.get("file_list", [])90 91 # Try downloading the model from this website.92 model_files = []93 for model_id, origin_file_path, local_dir in file_data:94 # Check if the file is downloaded.95 file_to_download = os.path.join(local_dir, os.path.basename(origin_file_path))96 if file_to_download in downloaded_files:97 continue98 # Download99 website_to_download_fn[website](model_id, origin_file_path, local_dir)100 if os.path.basename(origin_file_path) in os.listdir(local_dir):101 downloaded_files.append(file_to_download)102 model_files.append(file_to_download)103 104 # If the model is successfully downloaded, break.105 if len(model_files) > 0:106 if isinstance(model_metadata, dict) and "load_path" in model_metadata:107 model_files = model_metadata["load_path"]108 load_files.extend(model_files)109 break110 111 return load_files112 