#!/usr/bin/env python3 # -*- coding: utf-8 -*- """ 模型预下载脚本 统一从 ModelScope 预下载所有运行所需模型 """ import argparse import json from pathlib import Path from typing import Optional from modelscope.hub.snapshot_download import snapshot_download as ms_snapshot_download from app.core.config import settings from app.services.asr.model_capabilities import ( get_all_qwen_modelscope_assets, get_camplusplus_replacement_paths, get_download_modelscope_assets, ) def _get_qwen_modelscope_assets(): """Return all declared ModelScope Qwen assets for offline deployment.""" assets = get_all_qwen_modelscope_assets() if not assets: print("当前部署计划未启用 Qwen3-ASR,跳过 Qwen 模型下载") return [] print("离线部署模式:下载全部已声明 Qwen 模型(含 forced aligner)") return assets def _get_cache_path(model_id: str, source: str = "modelscope") -> Path: """获取模型缓存路径""" _ = source return Path(settings.MODELSCOPE_PATH) / model_id def check_model_exists(model_id: str, source: str = "modelscope") -> tuple[bool, str]: """检查模型是否已存在于本地缓存""" try: model_path = _get_cache_path(model_id, source) if model_path.exists() and model_path.is_dir(): if any(model_path.iterdir()): return True, str(model_path) except Exception: pass return False, "" def check_all_models() -> list[tuple[str, str, str, Optional[str]]]: """检查所有模型是否存在 Returns: 缺失的模型列表,每个元素为 (model_id, description, source, revision) """ missing = [] ms_assets = get_download_modelscope_assets() qwen_assets = _get_qwen_modelscope_assets() for asset in ms_assets: exists, _ = check_model_exists(asset.model_id, source="modelscope") if not exists: missing.append((asset.model_id, asset.description, "modelscope", asset.revision)) for asset in qwen_assets: exists, _ = check_model_exists(asset.model_id, source="modelscope") if not exists: missing.append((asset.model_id, asset.description, "modelscope", asset.revision)) return missing def fix_camplusplus_config() -> bool: """修复 CAM++ 配置文件,将模型ID替换为本地路径(用于离线环境) 修复 issue #15: 离线环境下 CAM++ 模型会尝试从 modelscope.cn 获取依赖模型配置 Returns: 是否修复成功 """ try: cache_dir = Path(settings.MODELSCOPE_PATH) config_file = cache_dir / "iic/speech_campplus_speaker-diarization_common/configuration.json" if not config_file.exists(): return False # 读取配置文件 with open(config_file, 'r', encoding='utf-8') as f: config = json.load(f) # 需要替换的模型ID -> 本地路径映射 replacements = get_camplusplus_replacement_paths(str(cache_dir)) # 检查是否需要修改 modified = False if "model" in config: for key in ["speaker_model", "change_locator", "vad_model"]: if key in config["model"]: old_value = config["model"][key] if old_value in replacements: new_value = replacements[old_value] # 检查本地路径是否存在 if Path(new_value).exists(): config["model"][key] = new_value modified = True # 写回配置文件 if modified: with open(config_file, 'w', encoding='utf-8') as f: json.dump(config, f, indent=4, ensure_ascii=False) return True return False except Exception as e: print(f"⚠️ 修复 CAM++ 配置文件失败: {e}") return False def download_models( auto_mode: bool = False, export_dir: Optional[str] = None, ) -> bool: """下载所有需要的模型 Args: auto_mode: 如果为True,表示自动模式(从start.py调用),会简化输出 export_dir: 如果指定,将下载的模型导出到该目录(用于离线部署) Returns: 是否全部下载成功 """ import shutil # 检查缺失的模型 missing = check_all_models() ms_assets = get_download_modelscope_assets() qwen_assets = _get_qwen_modelscope_assets() export_path = Path(export_dir) if export_dir else None if not missing: if not auto_mode: print("✅ 所有模型已存在,无需下载") if not export_path: return True ms_cache_dir = Path(settings.MODELSCOPE_CACHE) if auto_mode: print(f"📦 检测到 {len(missing)} 个模型需要下载...") else: print("=" * 60) print("Qwen3-ASR 模型预下载") print("=" * 60) print(f"ModelScope 缓存: {ms_cache_dir}") print(f"待下载模型: {len(missing)} 个") print("=" * 60) failed = [] downloaded = [] # 下载 ModelScope 模型 (Paraformer) ms_missing = [(mid, desc, rev) for mid, desc, src, rev in missing if src == "modelscope"] if ms_missing: if not auto_mode: print("\n📦 开始下载 ModelScope 模型 (Paraformer)...") print("-" * 60) for i, (model_id, desc, revision) in enumerate(ms_missing, 1): if not auto_mode: print(f"\n[{i}/{len(ms_missing)}] {desc}") print(f" 模型ID: {model_id}") if revision: print(f" 版本: {revision}") print(f" 📥 开始下载...", end="") try: local_dir = _get_cache_path(model_id, "modelscope") local_dir.parent.mkdir(parents=True, exist_ok=True) # 传递版本参数,如果指定了版本 if revision: path = ms_snapshot_download( model_id, revision=revision, cache_dir=str(ms_cache_dir), local_dir=str(local_dir), ) else: path = ms_snapshot_download( model_id, cache_dir=str(ms_cache_dir), local_dir=str(local_dir), ) if not auto_mode: print(f" ✅ 完成: {path}") downloaded.append((model_id, "modelscope", path)) except Exception as e: if not auto_mode: print(f" ❌ 失败: {e}") failed.append((model_id, str(e))) # 修复 CAM++ 配置文件(用于离线环境) if not auto_mode: print("\n🔧 修复 CAM++ 配置文件...") if fix_camplusplus_config(): if not auto_mode: print(" ✅ CAM++ 配置已修复(离线环境可用)") else: if not auto_mode: print(" ℹ️ 无需修复或配置文件不存在") # 导出模式:复制模型到扁平化的 models/ 根目录 if export_path and not failed: if not auto_mode: print(f"\n📦 导出模型到: {export_path}") # 收集所有需要导出的模型 all_models = [] for asset in ms_assets: all_models.append((asset.model_id, "modelscope")) for asset in qwen_assets: all_models.append((asset.model_id, "modelscope")) exported = 0 for model_entry in all_models: if len(model_entry) == 3: model_id, source, actual_model_id = model_entry else: model_id, source = model_entry actual_model_id = model_id cache_path = _get_cache_path(actual_model_id, source) if cache_path.exists(): rel_path = cache_path.relative_to(Path(settings.MODELSCOPE_PATH)) target_dir = export_path / rel_path target_dir.parent.mkdir(parents=True, exist_ok=True) if not auto_mode: print(f" 📂 {model_id}", end="") try: shutil.copytree(cache_path, target_dir, dirs_exist_ok=True) exported += 1 if not auto_mode: print(" ✅") except Exception as e: if not auto_mode: print(f" ❌ {e}") if not auto_mode: print(f"\n✅ 已导出 {exported} 个模型到 models/") if not auto_mode: print("\n" + "=" * 60) print("📊 下载统计:") print(f" ✅ 已下载: {len(downloaded)} 个") print(f" ❌ 失败: {len(failed)} 个") print("=" * 60) if failed: print(f"\n失败的模型:") for model_id, err in failed: print(f" - {model_id}: {err}") return False else: print("\n✅ 所有模型准备就绪!") print("=" * 60) return len(failed) == 0 def main() -> int: """CLI entrypoint for model download and export.""" parser = argparse.ArgumentParser(description="Download or export Qwen3-ASR models") parser.add_argument( "--export-dir", default=None, help="Optional export directory for offline deployment packaging", ) parser.add_argument( "--auto-mode", action="store_true", help="Reduce output for startup/bootstrap usage", ) args = parser.parse_args() success = download_models( auto_mode=args.auto_mode, export_dir=args.export_dir, ) return 0 if success else 1 if __name__ == "__main__": raise SystemExit(main())