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()