bert_base/scripts/download_models.py

163 lines
5.0 KiB
Python

import os
import sys
import shutil
import subprocess
from modelscope.hub.snapshot_download import snapshot_download
from dotenv import load_dotenv
load_dotenv()
def discover_models():
"""从环境变量中发现所有模型"""
models = {}
for key, value in os.environ.items():
if key.endswith('_MODEL_ID'):
prefix = key[:-9]
service_name = os.getenv(f'{prefix}_SERVICE_NAME', f'{prefix.lower()}_classifier')
service_folder = service_name.replace('_classifier', '')
model_type = os.getenv(f'{prefix}_TYPE', 'general')
models[prefix.lower()] = {
'prefix': prefix,
'model_id': value,
'port': os.getenv(f'{prefix}_PORT', '5002'),
'service_name': service_name,
'model_dir': os.getenv(f'{prefix}_MODEL_DIR', f'/app/services/{service_folder}/model'),
'service_folder': service_folder,
'model_type': model_type,
'cache_dir': os.path.join('/app/model_cache', prefix.lower())
}
return models
def find_model_files(cache_dir):
"""查找实际的模型文件位置"""
if os.path.exists(cache_dir):
for root, dirs, files in os.walk(cache_dir):
if any(f.endswith(('.safetensors', '.bin')) for f in files):
return root
return None
def download_model(name, config):
"""下载单个模型"""
print("=" * 50)
print(f"Processing {name} model...")
print(f"Model ID: {config['model_id']}")
print(f"Target: {config['model_dir']}")
print(f"Cache: {config['cache_dir']}")
print("=" * 50)
try:
# 确保目录存在
os.makedirs(config['model_dir'], exist_ok=True)
os.makedirs(config['cache_dir'], exist_ok=True)
# 检查是否已有模型文件
if os.path.exists(config['model_dir']):
files = [f for f in os.listdir(config['model_dir']) if f.endswith(('.safetensors', '.bin'))]
if files:
print(f"✓ Model already exists in {config['model_dir']} (found {len(files)} files)")
return True
# 登录 ModelScope
token = os.getenv('MODELSCOPE_TOKEN', '')
if token:
result = subprocess.run(
['modelscope', 'login', '--token', token],
capture_output=True,
text=True
)
if result.returncode == 0:
print("✓ ModelScope login successful")
else:
print(f"⚠️ ModelScope login failed: {result.stderr}")
# 下载模型
print(f"Downloading model to {config['cache_dir']}...")
snapshot_download(
model_id=config['model_id'],
cache_dir=config['cache_dir'],
revision="master",
ignore_file_pattern=[".git", ".gitattributes"]
)
# 查找实际模型文件
src_dir = find_model_files(config['cache_dir'])
if src_dir is None:
print(f"✗ Could not find model files in downloaded content")
return False
print(f"✓ Found model files in: {src_dir}")
# 复制文件到目标目录
print(f"Copying files to {config['model_dir']}...")
copied_count = 0
for file in os.listdir(src_dir):
src_file = os.path.join(src_dir, file)
if os.path.isfile(src_file):
dest_file = os.path.join(config['model_dir'], file)
shutil.copy2(src_file, dest_file)
copied_count += 1
print(f" Copied: {file}")
if copied_count > 0:
print(f"{name} model downloaded successfully! ({copied_count} files)")
return True
else:
print(f"✗ No files copied")
return False
except PermissionError as e:
print(f"✗ Permission error: {e}")
return False
except Exception as e:
print(f"✗ Failed to download {name} model: {e}")
return False
def main():
print("=" * 50)
print("Model Download Script")
print("=" * 50)
import pwd
current_user = pwd.getpwuid(os.getuid()).pw_name
print(f"Running as user: {current_user}")
print("=" * 50)
models = discover_models()
if not models:
print("⚠️ No models found in .env file!")
sys.exit(1)
print(f"\nFound {len(models)} model(s):")
for name, config in models.items():
print(f" {name}: {config['model_id']}")
print(f" Target: {config['model_dir']}")
print(f" Type: {config['model_type']}")
print("")
print("=" * 50)
print("Processing models...")
print("=" * 50)
success_count = 0
for name, config in models.items():
if download_model(name, config):
success_count += 1
print("")
print("=" * 50)
print(f"Processing completed: {success_count}/{len(models)} models")
print("=" * 50)
if success_count < len(models):
sys.exit(1)
if __name__ == "__main__":
main()