297 lines
9.7 KiB
Python
297 lines
9.7 KiB
Python
#!/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())
|